unstructured-platform-plugins 0.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- test/__init__.py +0 -0
- test/assets/__init__.py +0 -0
- test/assets/async_typed_dict_response.py +10 -0
- test/assets/dataclass_response.py +37 -0
- test/assets/empty_input_and_response.py +8 -0
- test/assets/hash_function.py +2 -0
- test/assets/improper_function.py +2 -0
- test/assets/pydantic_response_class_method.py +21 -0
- test/assets/simple_hash_class.py +6 -0
- test/assets/simple_hash_lambda.py +1 -0
- test/assets/simple_hash_value.py +1 -0
- test/assets/typed_dict_response.py +17 -0
- test/test_schema.py +660 -0
- test/test_utils.py +140 -0
- unstructured_platform_plugins/__init__.py +0 -0
- unstructured_platform_plugins/__version__.py +1 -0
- unstructured_platform_plugins/etl_uvicorn/__init__.py +0 -0
- unstructured_platform_plugins/etl_uvicorn/api_generator.py +173 -0
- unstructured_platform_plugins/etl_uvicorn/main.py +121 -0
- unstructured_platform_plugins/etl_uvicorn/utils.py +115 -0
- unstructured_platform_plugins/schema/__init__.py +0 -0
- unstructured_platform_plugins/schema/json_schema.py +344 -0
- unstructured_platform_plugins/schema/model.py +101 -0
- unstructured_platform_plugins/schema/usage.py +7 -0
- unstructured_platform_plugins/schema/utils.py +31 -0
- unstructured_platform_plugins/type_hints.py +106 -0
- unstructured_platform_plugins/validate_api.py +110 -0
- unstructured_platform_plugins-0.0.0.dist-info/LICENSE +201 -0
- unstructured_platform_plugins-0.0.0.dist-info/LICENSE.md +201 -0
- unstructured_platform_plugins-0.0.0.dist-info/METADATA +106 -0
- unstructured_platform_plugins-0.0.0.dist-info/RECORD +34 -0
- unstructured_platform_plugins-0.0.0.dist-info/WHEEL +5 -0
- unstructured_platform_plugins-0.0.0.dist-info/entry_points.txt +3 -0
- unstructured_platform_plugins-0.0.0.dist-info/top_level.txt +2 -0
test/test_utils.py
ADDED
|
@@ -0,0 +1,140 @@
|
|
|
1
|
+
import inspect
|
|
2
|
+
from dataclasses import dataclass, is_dataclass
|
|
3
|
+
from enum import Enum
|
|
4
|
+
|
|
5
|
+
import pytest
|
|
6
|
+
from pydantic import BaseModel
|
|
7
|
+
from unstructured.ingest.v2.interfaces import FileData
|
|
8
|
+
from uvicorn.importer import import_from_string
|
|
9
|
+
|
|
10
|
+
from unstructured_platform_plugins.etl_uvicorn import utils
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def test_get_func_simple():
|
|
14
|
+
instance = import_from_string("test.assets.typed_dict_response:sample_function")
|
|
15
|
+
func = utils.get_func(instance)
|
|
16
|
+
assert callable(func)
|
|
17
|
+
assert func.__name__ == "sample_function"
|
|
18
|
+
sig = inspect.signature(func)
|
|
19
|
+
response = sig.return_annotation
|
|
20
|
+
# Check for typed dict response
|
|
21
|
+
assert dict in inspect.getmro(response)
|
|
22
|
+
assert dict is not inspect.getmro(response)[0]
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def test_get_async_func():
|
|
26
|
+
instance = import_from_string("test.assets.async_typed_dict_response:async_sample_function")
|
|
27
|
+
func = utils.get_func(instance)
|
|
28
|
+
assert callable(func)
|
|
29
|
+
assert func.__name__ == "async_sample_function"
|
|
30
|
+
sig = inspect.signature(func)
|
|
31
|
+
response = sig.return_annotation
|
|
32
|
+
# Check for typed dict response
|
|
33
|
+
assert dict in inspect.getmro(response)
|
|
34
|
+
assert dict is not inspect.getmro(response)[0]
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def test_get_class_func():
|
|
38
|
+
instance = import_from_string("test.assets.pydantic_response_class_method:SampleClass")
|
|
39
|
+
func = utils.get_func(instance, method_name="sample_method")
|
|
40
|
+
assert callable(func)
|
|
41
|
+
assert func.__name__ == "sample_method"
|
|
42
|
+
sig = inspect.signature(func)
|
|
43
|
+
response = sig.return_annotation
|
|
44
|
+
# Check for pydantic model response
|
|
45
|
+
assert inspect.isclass(response)
|
|
46
|
+
assert issubclass(response, BaseModel)
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def test_get_func_dataclass_response():
|
|
50
|
+
instance = import_from_string("test.assets.dataclass_response:sample_function_with_path")
|
|
51
|
+
func = utils.get_func(instance)
|
|
52
|
+
assert callable(func)
|
|
53
|
+
assert func.__name__ == "sample_function_with_path"
|
|
54
|
+
sig = inspect.signature(func)
|
|
55
|
+
response = sig.return_annotation
|
|
56
|
+
# Check for dataclass response
|
|
57
|
+
assert is_dataclass(response)
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def test_get_plugin_id_value():
|
|
61
|
+
instance = import_from_string("test.assets.simple_hash_value:hash_value")
|
|
62
|
+
hash_value = utils.get_plugin_id(instance)
|
|
63
|
+
assert hash_value == "plugin_id_123"
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def test_get_plugin_id_lambda():
|
|
67
|
+
instance = import_from_string("test.assets.simple_hash_lambda:hash_lambda_fn")
|
|
68
|
+
hash_value = utils.get_plugin_id(instance)
|
|
69
|
+
assert hash_value == "plugin_id_hash_123"
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def test_get_plugin_id_function():
|
|
73
|
+
instance = import_from_string("test.assets.hash_function:get_hash")
|
|
74
|
+
hash_value = utils.get_plugin_id(instance)
|
|
75
|
+
assert hash_value == "plugin_id_fn_123"
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def test_get_plugin_id_class():
|
|
79
|
+
instance = import_from_string("test.assets.simple_hash_class:GetHash")
|
|
80
|
+
hash_value = utils.get_plugin_id(instance, method_name="my_hash")
|
|
81
|
+
assert hash_value == "plugin_id_class_123"
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def test_get_plugin_id_class_instance():
|
|
85
|
+
instance = import_from_string("test.assets.simple_hash_class:get_hash_class_instance")
|
|
86
|
+
hash_value = utils.get_plugin_id(instance, method_name="my_hash")
|
|
87
|
+
assert hash_value == "plugin_id_class_123"
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
@dataclass
|
|
91
|
+
class A:
|
|
92
|
+
b: int
|
|
93
|
+
c: float
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
class B(BaseModel):
|
|
97
|
+
d: bool
|
|
98
|
+
e: dict
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
class MyEnum(Enum):
|
|
102
|
+
VALUE = "value"
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def test_map_inputs():
|
|
106
|
+
def fn(a: A, b: B, c: MyEnum, d: list, e: FileData) -> None:
|
|
107
|
+
pass
|
|
108
|
+
|
|
109
|
+
file_data = FileData(
|
|
110
|
+
identifier="custom_file_data",
|
|
111
|
+
connector_type="mock_connector",
|
|
112
|
+
additional_metadata={"additional": "metadata"},
|
|
113
|
+
)
|
|
114
|
+
inputs = {
|
|
115
|
+
"a": {"b": 4, "c": 5.6},
|
|
116
|
+
"b": {"d": True, "e": {"key": "value"}},
|
|
117
|
+
"c": MyEnum.VALUE.value,
|
|
118
|
+
"d": [1, 2, 3],
|
|
119
|
+
"e": file_data.to_dict(),
|
|
120
|
+
}
|
|
121
|
+
|
|
122
|
+
mapped_inputs = utils.map_inputs(func=fn, raw_inputs=inputs)
|
|
123
|
+
expected = {
|
|
124
|
+
"a": A(b=4, c=5.6),
|
|
125
|
+
"b": B(d=True, e={"key": "value"}),
|
|
126
|
+
"c": MyEnum.VALUE.value,
|
|
127
|
+
"d": [1, 2, 3],
|
|
128
|
+
"e": file_data,
|
|
129
|
+
}
|
|
130
|
+
assert mapped_inputs == expected
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
def test_map_inputs_error():
|
|
134
|
+
def fn(a: FileData) -> None:
|
|
135
|
+
pass
|
|
136
|
+
|
|
137
|
+
inputs = {"a": {"not": "the", "right": "values"}}
|
|
138
|
+
|
|
139
|
+
with pytest.raises(KeyError):
|
|
140
|
+
utils.map_inputs(func=fn, raw_inputs=inputs)
|
|
File without changes
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "0.0.0" # pragma: no cover
|
|
File without changes
|
|
@@ -0,0 +1,173 @@
|
|
|
1
|
+
import asyncio
|
|
2
|
+
import hashlib
|
|
3
|
+
import inspect
|
|
4
|
+
import json
|
|
5
|
+
import logging
|
|
6
|
+
from typing import Any, Callable, Optional
|
|
7
|
+
|
|
8
|
+
from fastapi import FastAPI, status
|
|
9
|
+
from pydantic import BaseModel
|
|
10
|
+
from starlette.responses import RedirectResponse
|
|
11
|
+
from uvicorn.importer import import_from_string
|
|
12
|
+
|
|
13
|
+
from unstructured_platform_plugins.etl_uvicorn.utils import (
|
|
14
|
+
get_func,
|
|
15
|
+
get_input_schema,
|
|
16
|
+
get_output_sig,
|
|
17
|
+
get_plugin_id,
|
|
18
|
+
get_schema_dict,
|
|
19
|
+
map_inputs,
|
|
20
|
+
)
|
|
21
|
+
from unstructured_platform_plugins.schema.json_schema import (
|
|
22
|
+
schema_to_base_model,
|
|
23
|
+
)
|
|
24
|
+
from unstructured_platform_plugins.schema.usage import UsageData
|
|
25
|
+
|
|
26
|
+
logger = logging.getLogger("uvicorn.error")
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
async def invoke_func(func: Callable, kwargs: Optional[dict[str, Any]] = None) -> Any:
|
|
30
|
+
kwargs = kwargs or {}
|
|
31
|
+
if inspect.iscoroutinefunction(func):
|
|
32
|
+
return await func(**kwargs)
|
|
33
|
+
else:
|
|
34
|
+
return func(**kwargs)
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def check_precheck_func(precheck_func: Callable):
|
|
38
|
+
sig = inspect.signature(precheck_func)
|
|
39
|
+
inputs = sig.parameters.values()
|
|
40
|
+
outputs = sig.return_annotation
|
|
41
|
+
if len(inputs) == 1:
|
|
42
|
+
i = inputs[0]
|
|
43
|
+
if i.name != "usage" or i.annotation is list:
|
|
44
|
+
raise ValueError("the only input available for precheck is usage which must be a list")
|
|
45
|
+
if outputs not in [None, sig.empty]:
|
|
46
|
+
raise ValueError(f"no output should exist for precheck function, found: {outputs}")
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def generate_fast_api(
|
|
50
|
+
app: str,
|
|
51
|
+
method_name: Optional[str] = None,
|
|
52
|
+
id_str: Optional[str] = None,
|
|
53
|
+
id_method: Optional[str] = None,
|
|
54
|
+
precheck_str: Optional[str] = None,
|
|
55
|
+
precheck_method: Optional[str] = None,
|
|
56
|
+
) -> FastAPI:
|
|
57
|
+
instance = import_from_string(app)
|
|
58
|
+
func = get_func(instance, method_name)
|
|
59
|
+
if id_str:
|
|
60
|
+
id_ref = import_from_string(id_str)
|
|
61
|
+
plugin_id = get_plugin_id(instance=id_ref, method_name=id_method)
|
|
62
|
+
else:
|
|
63
|
+
plugin_id = hashlib.sha256(
|
|
64
|
+
json.dumps(get_schema_dict(func), sort_keys=True).encode()
|
|
65
|
+
).hexdigest()[:32]
|
|
66
|
+
|
|
67
|
+
precheck_func = None
|
|
68
|
+
if precheck_str:
|
|
69
|
+
precheck_instance = import_from_string(precheck_str)
|
|
70
|
+
precheck_func = get_func(precheck_instance, precheck_method)
|
|
71
|
+
elif precheck_method:
|
|
72
|
+
precheck_func = get_func(instance, precheck_method)
|
|
73
|
+
if precheck_func is not None:
|
|
74
|
+
check_precheck_func(precheck_func=precheck_func)
|
|
75
|
+
|
|
76
|
+
logger.debug(f"set static id response to: {plugin_id}")
|
|
77
|
+
|
|
78
|
+
fastapi_app = FastAPI()
|
|
79
|
+
|
|
80
|
+
response_type = get_output_sig(func)
|
|
81
|
+
|
|
82
|
+
class InvokeResponse(BaseModel):
|
|
83
|
+
usage: list[UsageData]
|
|
84
|
+
status_code: int
|
|
85
|
+
status_code_text: Optional[str] = None
|
|
86
|
+
output: Optional[response_type] = None
|
|
87
|
+
|
|
88
|
+
input_schema = get_input_schema(func, omit=["usage"])
|
|
89
|
+
input_schema_model = schema_to_base_model(input_schema)
|
|
90
|
+
|
|
91
|
+
logging.getLogger("etl_uvicorn.fastapi")
|
|
92
|
+
|
|
93
|
+
async def wrap_fn(func: Callable, kwargs: Optional[dict[str, Any]] = None) -> InvokeResponse:
|
|
94
|
+
usage: list[UsageData] = []
|
|
95
|
+
request_dict = kwargs if kwargs else {}
|
|
96
|
+
if "usage" in inspect.signature(func).parameters:
|
|
97
|
+
request_dict["usage"] = usage
|
|
98
|
+
else:
|
|
99
|
+
logger.warning("usage data not an expected parameter, omitting")
|
|
100
|
+
try:
|
|
101
|
+
output = await invoke_func(func=func, kwargs=request_dict)
|
|
102
|
+
return InvokeResponse(usage=usage, status_code=status.HTTP_200_OK, output=output)
|
|
103
|
+
except Exception as invoke_error:
|
|
104
|
+
logger.error(f"failed to invoke plugin: {invoke_error}", exc_info=True)
|
|
105
|
+
return InvokeResponse(
|
|
106
|
+
usage=usage,
|
|
107
|
+
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
|
108
|
+
status_code_text=f"failed to invoke plugin: "
|
|
109
|
+
f"[{invoke_error.__class__.__name__}] {invoke_error}",
|
|
110
|
+
)
|
|
111
|
+
|
|
112
|
+
if input_schema_model.model_fields:
|
|
113
|
+
|
|
114
|
+
@fastapi_app.post("/invoke", response_model=InvokeResponse)
|
|
115
|
+
async def run_job(request: input_schema_model) -> InvokeResponse:
|
|
116
|
+
logger.debug(f"invoking function {func} with input: {request.model_dump()}")
|
|
117
|
+
# Create dictionary from pydantic model while preserving underlying types
|
|
118
|
+
request_dict = {f: getattr(request, f) for f in request.model_fields}
|
|
119
|
+
map_inputs(func=func, raw_inputs=request_dict)
|
|
120
|
+
logger.debug(f"passing inputs to function: {request_dict}")
|
|
121
|
+
return await wrap_fn(func=func, kwargs=request_dict)
|
|
122
|
+
|
|
123
|
+
else:
|
|
124
|
+
|
|
125
|
+
@fastapi_app.post("/invoke", response_model=response_type)
|
|
126
|
+
async def run_job() -> response_type:
|
|
127
|
+
logger.debug(f"invoking function without inputs: {func}")
|
|
128
|
+
return await wrap_fn(
|
|
129
|
+
func=func,
|
|
130
|
+
)
|
|
131
|
+
|
|
132
|
+
class SchemaOutputResponse(BaseModel):
|
|
133
|
+
inputs: dict[str, Any]
|
|
134
|
+
outputs: dict[str, Any]
|
|
135
|
+
|
|
136
|
+
@fastapi_app.get("/", include_in_schema=False)
|
|
137
|
+
async def docs_redirect():
|
|
138
|
+
return RedirectResponse("/docs")
|
|
139
|
+
|
|
140
|
+
class InvokePrecheckResponse(BaseModel):
|
|
141
|
+
usage: list[UsageData]
|
|
142
|
+
status_code: int
|
|
143
|
+
status_code_text: Optional[str] = None
|
|
144
|
+
|
|
145
|
+
@fastapi_app.get("/schema")
|
|
146
|
+
async def get_schema() -> SchemaOutputResponse:
|
|
147
|
+
schema = get_schema_dict(func)
|
|
148
|
+
resp = SchemaOutputResponse(inputs=schema["inputs"], outputs=schema["outputs"])
|
|
149
|
+
return resp
|
|
150
|
+
|
|
151
|
+
@fastapi_app.get("/precheck")
|
|
152
|
+
async def run_precheck() -> InvokePrecheckResponse:
|
|
153
|
+
if precheck_func:
|
|
154
|
+
fn_response = await wrap_fn(func=precheck_func)
|
|
155
|
+
return InvokePrecheckResponse(
|
|
156
|
+
status_code=fn_response.status_code,
|
|
157
|
+
status_code_text=fn_response.status_code_text,
|
|
158
|
+
usage=fn_response.usage,
|
|
159
|
+
)
|
|
160
|
+
else:
|
|
161
|
+
return InvokePrecheckResponse(status_code=status.HTTP_200_OK, usage=[])
|
|
162
|
+
|
|
163
|
+
@fastapi_app.get("/id")
|
|
164
|
+
async def get_id() -> str:
|
|
165
|
+
return plugin_id
|
|
166
|
+
|
|
167
|
+
# Run initial schema validation
|
|
168
|
+
try:
|
|
169
|
+
asyncio.run(get_schema())
|
|
170
|
+
except TypeError as e:
|
|
171
|
+
raise TypeError(f"failed to validate function schema: {e}") from e
|
|
172
|
+
|
|
173
|
+
return fastapi_app
|
|
@@ -0,0 +1,121 @@
|
|
|
1
|
+
from dataclasses import dataclass, field
|
|
2
|
+
from typing import IO, Any, Optional
|
|
3
|
+
|
|
4
|
+
import click
|
|
5
|
+
from uvicorn.config import LOGGING_CONFIG, Config, RawConfigParser
|
|
6
|
+
from uvicorn.main import main, run
|
|
7
|
+
|
|
8
|
+
from unstructured_platform_plugins.etl_uvicorn.api_generator import generate_fast_api
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
@dataclass
|
|
12
|
+
class CustomConfig:
|
|
13
|
+
log_config: dict[str, Any] | str | RawConfigParser | IO[Any] | None = field(
|
|
14
|
+
default_factory=lambda: LOGGING_CONFIG
|
|
15
|
+
)
|
|
16
|
+
log_level: str | int | None = None
|
|
17
|
+
use_colors: bool | None = None
|
|
18
|
+
access_log: bool = True
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
CustomConfig.configure_logging = Config.configure_logging
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def get_command() -> click.Command:
|
|
25
|
+
@click.command(context_settings={"auto_envvar_prefix": "UVICORN"})
|
|
26
|
+
def api_wrapper(
|
|
27
|
+
app: str,
|
|
28
|
+
log_config: str,
|
|
29
|
+
reload_dirs: list[str],
|
|
30
|
+
reload_includes: list[str],
|
|
31
|
+
reload_excludes: list[str],
|
|
32
|
+
headers: list[str],
|
|
33
|
+
method_name: Optional[str] = None,
|
|
34
|
+
plugin_id: Optional[str] = None,
|
|
35
|
+
plugin_id_method: Optional[str] = None,
|
|
36
|
+
precheck_app: Optional[str] = None,
|
|
37
|
+
precheck_app_method: Optional[str] = None,
|
|
38
|
+
**kwargs,
|
|
39
|
+
):
|
|
40
|
+
# Make sure logging is configured before the call to run() so any setup has the same format
|
|
41
|
+
config = CustomConfig(
|
|
42
|
+
log_config=LOGGING_CONFIG if log_config is None else log_config,
|
|
43
|
+
log_level=kwargs["log_level"],
|
|
44
|
+
use_colors=kwargs["use_colors"],
|
|
45
|
+
access_log=kwargs["access_log"],
|
|
46
|
+
)
|
|
47
|
+
config.configure_logging()
|
|
48
|
+
fastapi_app = generate_fast_api(
|
|
49
|
+
app=app,
|
|
50
|
+
method_name=method_name,
|
|
51
|
+
id_str=plugin_id,
|
|
52
|
+
id_method=plugin_id_method,
|
|
53
|
+
precheck_str=precheck_app,
|
|
54
|
+
precheck_method=precheck_app_method,
|
|
55
|
+
)
|
|
56
|
+
# Explicitly map values that are manipulated in the original
|
|
57
|
+
# call to run(), preventing **kwargs reference
|
|
58
|
+
run(
|
|
59
|
+
fastapi_app,
|
|
60
|
+
log_config=LOGGING_CONFIG if log_config is None else log_config,
|
|
61
|
+
reload_dirs=reload_dirs or None,
|
|
62
|
+
reload_includes=reload_includes or None,
|
|
63
|
+
reload_excludes=reload_excludes or None,
|
|
64
|
+
headers=[header.split(":", 1) for header in headers], # type: ignore[misc]
|
|
65
|
+
**kwargs,
|
|
66
|
+
)
|
|
67
|
+
|
|
68
|
+
cmd = api_wrapper
|
|
69
|
+
cmd.params = main.params
|
|
70
|
+
cmd.params.extend(
|
|
71
|
+
[
|
|
72
|
+
click.Option(
|
|
73
|
+
["--method-name"],
|
|
74
|
+
required=False,
|
|
75
|
+
type=str,
|
|
76
|
+
default=None,
|
|
77
|
+
help="If passed in instance is a class, what method to wrap. "
|
|
78
|
+
"Will fall back to __call__ if none is provided.",
|
|
79
|
+
),
|
|
80
|
+
click.Option(
|
|
81
|
+
["--plugin-id"],
|
|
82
|
+
required=False,
|
|
83
|
+
type=str,
|
|
84
|
+
default=None,
|
|
85
|
+
help="Reference to either a value or function to get "
|
|
86
|
+
"the plugin id once instantiated",
|
|
87
|
+
),
|
|
88
|
+
click.Option(
|
|
89
|
+
["--plugin-id-method"],
|
|
90
|
+
required=False,
|
|
91
|
+
type=str,
|
|
92
|
+
default=None,
|
|
93
|
+
help="If plugin id reference is a class, what method to wrap. "
|
|
94
|
+
"Will fall back to __call__ if none is provided.",
|
|
95
|
+
),
|
|
96
|
+
click.Option(
|
|
97
|
+
["--precheck-app"],
|
|
98
|
+
required=False,
|
|
99
|
+
type=str,
|
|
100
|
+
default=None,
|
|
101
|
+
help="If provided, must point to code to run for precheck",
|
|
102
|
+
),
|
|
103
|
+
click.Option(
|
|
104
|
+
["--precheck-app-method"],
|
|
105
|
+
required=False,
|
|
106
|
+
type=str,
|
|
107
|
+
default=None,
|
|
108
|
+
help="If provided, points to a method to call on a class. "
|
|
109
|
+
"If precheck-app not provided, assumes method "
|
|
110
|
+
"lives on main class passes in.",
|
|
111
|
+
),
|
|
112
|
+
]
|
|
113
|
+
)
|
|
114
|
+
return cmd
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
cmd = get_command()
|
|
118
|
+
|
|
119
|
+
if __name__ == "__main__":
|
|
120
|
+
cmd = get_command()
|
|
121
|
+
cmd()
|
|
@@ -0,0 +1,115 @@
|
|
|
1
|
+
import inspect
|
|
2
|
+
from dataclasses import is_dataclass
|
|
3
|
+
from enum import EnumMeta
|
|
4
|
+
from types import GenericAlias, NoneType
|
|
5
|
+
from typing import Any, Callable, Optional
|
|
6
|
+
|
|
7
|
+
from dataclasses_json import DataClassJsonMixin
|
|
8
|
+
from pydantic import BaseModel
|
|
9
|
+
|
|
10
|
+
from unstructured_platform_plugins.schema.json_schema import (
|
|
11
|
+
parameters_to_json_schema,
|
|
12
|
+
response_to_json_schema,
|
|
13
|
+
)
|
|
14
|
+
from unstructured_platform_plugins.schema.utils import get_typed_parameters
|
|
15
|
+
from unstructured_platform_plugins.type_hints import get_type_hints
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def get_func(instance: Any, method_name: Optional[str] = None) -> Callable:
|
|
19
|
+
method_name = method_name or "__call__"
|
|
20
|
+
if inspect.isfunction(instance):
|
|
21
|
+
return instance
|
|
22
|
+
elif inspect.isclass(instance):
|
|
23
|
+
i = instance()
|
|
24
|
+
return getattr(i, method_name)
|
|
25
|
+
elif isinstance(instance, object) and hasattr(instance, method_name):
|
|
26
|
+
func = getattr(instance, method_name)
|
|
27
|
+
if inspect.ismethod(func):
|
|
28
|
+
return func
|
|
29
|
+
raise ValueError(f"type of instance not recognized: {type(instance)}")
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def get_plugin_id(instance: Any, method_name: Optional[str] = None) -> str:
|
|
33
|
+
method_name = method_name or "__call__"
|
|
34
|
+
ref_id = None
|
|
35
|
+
if inspect.isfunction(instance):
|
|
36
|
+
ref_id = instance()
|
|
37
|
+
elif inspect.isclass(instance):
|
|
38
|
+
i = instance()
|
|
39
|
+
method_name = method_name or "__call__"
|
|
40
|
+
fn = getattr(i, method_name)
|
|
41
|
+
ref_id = fn()
|
|
42
|
+
elif isinstance(instance, object) and hasattr(instance, method_name):
|
|
43
|
+
func = getattr(instance, method_name)
|
|
44
|
+
if inspect.ismethod(func):
|
|
45
|
+
ref_id = func()
|
|
46
|
+
else:
|
|
47
|
+
ref_id = instance
|
|
48
|
+
if not ref_id:
|
|
49
|
+
raise ValueError(f"id could not be parsed from instance {instance}")
|
|
50
|
+
ref_id = str(ref_id)
|
|
51
|
+
if not ref_id.isidentifier():
|
|
52
|
+
raise ValueError(f"'{ref_id}' is not a valid identifier")
|
|
53
|
+
return ref_id
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def get_input_schema(func: Callable, omit: Optional[list[str]] = None) -> dict:
|
|
57
|
+
|
|
58
|
+
parameters = get_typed_parameters(func)
|
|
59
|
+
if omit:
|
|
60
|
+
parameters = [p for p in parameters if p.name not in omit]
|
|
61
|
+
return parameters_to_json_schema(parameters)
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def get_output_sig(func: Callable) -> Optional[Any]:
|
|
65
|
+
inspect.signature(func)
|
|
66
|
+
type_hints = get_type_hints(func)
|
|
67
|
+
return_typing = type_hints.get("return")
|
|
68
|
+
outputs = return_typing if return_typing is not NoneType else None
|
|
69
|
+
return outputs
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def get_output_schema(func: Callable) -> dict:
|
|
73
|
+
return response_to_json_schema(get_output_sig(func))
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def get_schema_dict(func, omit: list[str] = ["usage"]) -> dict:
|
|
77
|
+
return {
|
|
78
|
+
"inputs": get_input_schema(func, omit=omit),
|
|
79
|
+
"outputs": get_output_schema(func),
|
|
80
|
+
}
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def map_inputs(func: Callable, raw_inputs: dict[str, Any]) -> dict[str, Any]:
|
|
84
|
+
# deserializes the raw dictionary coming in from the api into the underlying data
|
|
85
|
+
# types expected by the function when being invoked
|
|
86
|
+
raw_inputs = raw_inputs.copy()
|
|
87
|
+
type_info = get_type_hints(func)
|
|
88
|
+
type_info.pop("return", None)
|
|
89
|
+
for field_name, type_data in type_info.items():
|
|
90
|
+
if field_name not in raw_inputs:
|
|
91
|
+
continue
|
|
92
|
+
field_value = raw_inputs[field_name]
|
|
93
|
+
try:
|
|
94
|
+
if (
|
|
95
|
+
inspect.isclass(type_data)
|
|
96
|
+
and issubclass(type_data, DataClassJsonMixin)
|
|
97
|
+
and isinstance(field_value, dict)
|
|
98
|
+
):
|
|
99
|
+
raw_inputs[field_name] = type_data.from_dict(raw_inputs[field_name])
|
|
100
|
+
elif is_dataclass(type_data) and isinstance(field_value, dict):
|
|
101
|
+
raw_inputs[field_name] = type_data(**raw_inputs[field_name])
|
|
102
|
+
elif isinstance(type_data, EnumMeta):
|
|
103
|
+
raw_inputs[field_name] = raw_inputs[field_name]
|
|
104
|
+
elif (
|
|
105
|
+
inspect.isclass(type_data)
|
|
106
|
+
and not isinstance(type_data, GenericAlias)
|
|
107
|
+
and issubclass(type_data, BaseModel)
|
|
108
|
+
):
|
|
109
|
+
raw_inputs[field_name] = type_data.model_validate(raw_inputs[field_name])
|
|
110
|
+
except Exception as e:
|
|
111
|
+
exception_type = type(e)
|
|
112
|
+
raise exception_type(
|
|
113
|
+
f"failed to map input for field {field_name}: {field_value}"
|
|
114
|
+
) from e
|
|
115
|
+
return raw_inputs
|
|
File without changes
|