refactor: move logic out from host.py
This commit is contained in:
@@ -1,8 +1,9 @@
|
|||||||
from .api.host import (
|
from .api.host import ExtensionHostAPI
|
||||||
ExtensionAPIMethod,
|
from .api.models import ExtensionAPIMethod, ExtensionAPIMethodExport
|
||||||
ExtensionHostAPI,
|
from .api.registry import (
|
||||||
extension_api_contract,
|
extension_api_contract,
|
||||||
extension_api_method,
|
extension_api_method,
|
||||||
|
extension_api_permission_ids,
|
||||||
get_extension_api_method,
|
get_extension_api_method,
|
||||||
list_extension_api_methods,
|
list_extension_api_methods,
|
||||||
)
|
)
|
||||||
@@ -12,10 +13,12 @@ from .wasm.loader import WasmExtension
|
|||||||
__all__ = [
|
__all__ = [
|
||||||
"ExtensionAPIHost",
|
"ExtensionAPIHost",
|
||||||
"ExtensionAPIMethod",
|
"ExtensionAPIMethod",
|
||||||
|
"ExtensionAPIMethodExport",
|
||||||
"ExtensionHostAPI",
|
"ExtensionHostAPI",
|
||||||
"WasmExtension",
|
"WasmExtension",
|
||||||
"extension_api_contract",
|
"extension_api_contract",
|
||||||
"extension_api_method",
|
"extension_api_method",
|
||||||
|
"extension_api_permission_ids",
|
||||||
"get_extension_api_method",
|
"get_extension_api_method",
|
||||||
"list_extension_api_methods",
|
"list_extension_api_methods",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -1,8 +1,9 @@
|
|||||||
from .host import (
|
from .host import ExtensionHostAPI
|
||||||
ExtensionAPIMethod,
|
from .models import ExtensionAPIMethod, ExtensionAPIMethodExport
|
||||||
ExtensionHostAPI,
|
from .registry import (
|
||||||
extension_api_contract,
|
extension_api_contract,
|
||||||
extension_api_method,
|
extension_api_method,
|
||||||
|
extension_api_permission_ids,
|
||||||
get_extension_api_method,
|
get_extension_api_method,
|
||||||
list_extension_api_methods,
|
list_extension_api_methods,
|
||||||
)
|
)
|
||||||
@@ -11,9 +12,11 @@ from .runtime import ExtensionAPIHost
|
|||||||
__all__ = [
|
__all__ = [
|
||||||
"ExtensionAPIHost",
|
"ExtensionAPIHost",
|
||||||
"ExtensionAPIMethod",
|
"ExtensionAPIMethod",
|
||||||
|
"ExtensionAPIMethodExport",
|
||||||
"ExtensionHostAPI",
|
"ExtensionHostAPI",
|
||||||
"extension_api_contract",
|
"extension_api_contract",
|
||||||
"extension_api_method",
|
"extension_api_method",
|
||||||
|
"extension_api_permission_ids",
|
||||||
"get_extension_api_method",
|
"get_extension_api_method",
|
||||||
"list_extension_api_methods",
|
"list_extension_api_methods",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -1,16 +1,11 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import inspect
|
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import secrets
|
import secrets
|
||||||
import time
|
import time
|
||||||
from collections.abc import Awaitable, Callable, Iterable, Mapping
|
from collections.abc import Iterable, Mapping
|
||||||
from dataclasses import dataclass
|
from typing import Any
|
||||||
from functools import wraps
|
|
||||||
from typing import Any, TypeVar, cast, get_type_hints
|
|
||||||
|
|
||||||
from pydantic import BaseModel
|
|
||||||
|
|
||||||
from lnbits.helpers import sha256s
|
from lnbits.helpers import sha256s
|
||||||
|
|
||||||
@@ -49,93 +44,10 @@ from .models import (
|
|||||||
WalletBalanceRequest,
|
WalletBalanceRequest,
|
||||||
WalletBalanceResponse,
|
WalletBalanceResponse,
|
||||||
)
|
)
|
||||||
|
from .registry import extension_api_method
|
||||||
|
|
||||||
logger = logging.getLogger("lnbits.extensions")
|
logger = logging.getLogger("lnbits.extensions")
|
||||||
|
|
||||||
_EXTENSION_API_METHOD_ATTR = "__lnbits_extension_api_method__"
|
|
||||||
_EXTENSION_RUNTIME_PERMISSION_IDS = {"ui.camera.scan_qr"}
|
|
||||||
_RequestModel = TypeVar("_RequestModel", bound=BaseModel)
|
|
||||||
_ResponseModel = TypeVar("_ResponseModel", bound=BaseModel)
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
|
||||||
class ExtensionAPIMethodExport:
|
|
||||||
method_id: str
|
|
||||||
namespace: str
|
|
||||||
name: str
|
|
||||||
host_interface: str
|
|
||||||
host_name: str
|
|
||||||
sdk_name: str
|
|
||||||
description: str
|
|
||||||
required_permission: str | None = None
|
|
||||||
require_auth: bool = True
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
|
||||||
class ExtensionAPIMethod:
|
|
||||||
method_id: str
|
|
||||||
namespace: str
|
|
||||||
name: str
|
|
||||||
python_name: str
|
|
||||||
host_interface: str
|
|
||||||
host_name: str
|
|
||||||
sdk_name: str
|
|
||||||
description: str
|
|
||||||
request_model: type[BaseModel]
|
|
||||||
response_model: type[BaseModel]
|
|
||||||
required_permission: str | None = None
|
|
||||||
require_auth: bool = True
|
|
||||||
|
|
||||||
@property
|
|
||||||
def sdk_qualified_name(self) -> str:
|
|
||||||
return f"{self.namespace}.{self.sdk_name}"
|
|
||||||
|
|
||||||
|
|
||||||
def extension_api_method(
|
|
||||||
*,
|
|
||||||
method_id: str,
|
|
||||||
namespace: str,
|
|
||||||
name: str,
|
|
||||||
host_name: str,
|
|
||||||
sdk_name: str,
|
|
||||||
description: str,
|
|
||||||
host_interface: str = "host",
|
|
||||||
required_permission: str | None = None,
|
|
||||||
require_auth: bool = True,
|
|
||||||
) -> Callable[
|
|
||||||
[Callable[[Any, _RequestModel], Awaitable[_ResponseModel]]],
|
|
||||||
Callable[[Any, _RequestModel], Awaitable[_ResponseModel]],
|
|
||||||
]:
|
|
||||||
export = ExtensionAPIMethodExport(
|
|
||||||
method_id=method_id,
|
|
||||||
namespace=namespace,
|
|
||||||
name=name,
|
|
||||||
host_interface=host_interface,
|
|
||||||
host_name=host_name,
|
|
||||||
sdk_name=sdk_name,
|
|
||||||
description=description,
|
|
||||||
required_permission=required_permission,
|
|
||||||
require_auth=require_auth,
|
|
||||||
)
|
|
||||||
|
|
||||||
def decorator(
|
|
||||||
function: Callable[[Any, _RequestModel], Awaitable[_ResponseModel]],
|
|
||||||
) -> Callable[[Any, _RequestModel], Awaitable[_ResponseModel]]:
|
|
||||||
@wraps(function)
|
|
||||||
async def wrapper(self: Any, request: _RequestModel) -> _ResponseModel:
|
|
||||||
api = getattr(self, "api", self)
|
|
||||||
if require_auth and not api.has_authenticated_context():
|
|
||||||
raise PermissionError(
|
|
||||||
f"Extension API method '{method_id}' requires authentication."
|
|
||||||
)
|
|
||||||
api.require_permission(required_permission)
|
|
||||||
return await function(self, request)
|
|
||||||
|
|
||||||
setattr(wrapper, _EXTENSION_API_METHOD_ATTR, export)
|
|
||||||
return wrapper
|
|
||||||
|
|
||||||
return decorator
|
|
||||||
|
|
||||||
|
|
||||||
class ExtensionHostAPI:
|
class ExtensionHostAPI:
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -154,34 +66,10 @@ class ExtensionHostAPI:
|
|||||||
self.access_token = access_token
|
self.access_token = access_token
|
||||||
self.context = context
|
self.context = context
|
||||||
self.owner_id = sha256s(user_id) if user_id else owner_id
|
self.owner_id = sha256s(user_id) if user_id else owner_id
|
||||||
self._uuid = secrets.token_urlsafe(12).replace("-", "_")
|
|
||||||
from .utils import ExtensionAPIUtils
|
from .utils import ExtensionAPIUtils
|
||||||
|
|
||||||
self.utils = ExtensionAPIUtils(self)
|
self.utils = ExtensionAPIUtils(self)
|
||||||
|
|
||||||
def __repr__(self) -> str:
|
|
||||||
return (
|
|
||||||
"ExtensionHostAPI("
|
|
||||||
f"extension_id={self.extension_id!r}, "
|
|
||||||
f"context={self.context!r}, "
|
|
||||||
f"_uuid={self._uuid!r}"
|
|
||||||
")"
|
|
||||||
)
|
|
||||||
|
|
||||||
def require_permission(self, permission: str | None) -> None:
|
|
||||||
if permission and permission not in self.permissions:
|
|
||||||
raise PermissionError(
|
|
||||||
f"Extension '{self.extension_id}' is missing permission '{permission}'."
|
|
||||||
)
|
|
||||||
|
|
||||||
def has_authenticated_context(self) -> bool:
|
|
||||||
return bool(self.user_id) or self.context == "event"
|
|
||||||
|
|
||||||
def _require_owner_id(self) -> str:
|
|
||||||
if not self.owner_id:
|
|
||||||
raise PermissionError("Extension API method requires an owner context.")
|
|
||||||
return self.owner_id
|
|
||||||
|
|
||||||
@extension_api_method(
|
@extension_api_method(
|
||||||
method_id="storage.get",
|
method_id="storage.get",
|
||||||
namespace="storage",
|
namespace="storage",
|
||||||
@@ -223,6 +111,7 @@ class ExtensionHostAPI:
|
|||||||
for field_name, value in row.items()
|
for field_name, value in row.items()
|
||||||
if field_name in public_fields
|
if field_name in public_fields
|
||||||
}
|
}
|
||||||
|
# todo: check public fields filtering
|
||||||
return StorageGetResponse(data_json=json.dumps(public_row))
|
return StorageGetResponse(data_json=json.dumps(public_row))
|
||||||
|
|
||||||
@extension_api_method(
|
@extension_api_method(
|
||||||
@@ -647,131 +536,25 @@ class ExtensionHostAPI:
|
|||||||
)
|
)
|
||||||
return table, wallet_field
|
return table, wallet_field
|
||||||
|
|
||||||
|
def require_permission(self, permission: str | None) -> None:
|
||||||
def list_extension_api_methods(
|
if permission and permission not in self.permissions:
|
||||||
api_cls: type[ExtensionHostAPI] = ExtensionHostAPI,
|
raise PermissionError(
|
||||||
) -> list[ExtensionAPIMethod]:
|
f"Extension '{self.extension_id}' is missing permission '{permission}'."
|
||||||
methods: list[ExtensionAPIMethod] = []
|
|
||||||
|
|
||||||
for prefix, method_cls in _extension_api_method_sources(api_cls):
|
|
||||||
for python_name, function in inspect.getmembers(method_cls, inspect.isfunction):
|
|
||||||
export = getattr(function, _EXTENSION_API_METHOD_ATTR, None)
|
|
||||||
if not export:
|
|
||||||
continue
|
|
||||||
|
|
||||||
request_model, response_model = _get_method_models(function)
|
|
||||||
methods.append(
|
|
||||||
ExtensionAPIMethod(
|
|
||||||
method_id=export.method_id,
|
|
||||||
namespace=export.namespace,
|
|
||||||
name=export.name,
|
|
||||||
python_name=f"{prefix}.{python_name}" if prefix else python_name,
|
|
||||||
host_interface=export.host_interface,
|
|
||||||
host_name=export.host_name,
|
|
||||||
sdk_name=export.sdk_name,
|
|
||||||
description=export.description,
|
|
||||||
request_model=request_model,
|
|
||||||
response_model=response_model,
|
|
||||||
required_permission=export.required_permission,
|
|
||||||
require_auth=export.require_auth,
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
return sorted(methods, key=lambda method: method.method_id)
|
def has_authenticated_context(self) -> bool:
|
||||||
|
return bool(self.user_id) or self.context == "event"
|
||||||
|
|
||||||
|
def _require_owner_id(self) -> str:
|
||||||
|
if not self.owner_id:
|
||||||
|
raise PermissionError("Extension API method requires an owner context.")
|
||||||
|
return self.owner_id
|
||||||
|
|
||||||
def _extension_api_method_sources(
|
def __repr__(self) -> str:
|
||||||
api_cls: type[ExtensionHostAPI],
|
return (
|
||||||
) -> list[tuple[str, type[Any]]]:
|
"ExtensionHostAPI("
|
||||||
sources: list[tuple[str, type[Any]]] = [("", api_cls)]
|
f"extension_id={self.extension_id!r}, "
|
||||||
if issubclass(api_cls, ExtensionHostAPI):
|
f"context={self.context!r}, "
|
||||||
from .utils import extension_api_utils_method_classes
|
f"owner_id={self.owner_id!r}"
|
||||||
|
")"
|
||||||
sources.extend(extension_api_utils_method_classes().items())
|
|
||||||
return sources
|
|
||||||
|
|
||||||
|
|
||||||
def extension_api_permission_ids(
|
|
||||||
api_cls: type[ExtensionHostAPI] = ExtensionHostAPI,
|
|
||||||
) -> set[str]:
|
|
||||||
permissions = {
|
|
||||||
method.required_permission
|
|
||||||
for method in list_extension_api_methods(api_cls)
|
|
||||||
if method.required_permission
|
|
||||||
}
|
|
||||||
permissions.update(_EXTENSION_RUNTIME_PERMISSION_IDS)
|
|
||||||
return permissions
|
|
||||||
|
|
||||||
|
|
||||||
def get_extension_api_method(
|
|
||||||
method_id: str,
|
|
||||||
api_cls: type[ExtensionHostAPI] = ExtensionHostAPI,
|
|
||||||
) -> ExtensionAPIMethod:
|
|
||||||
for method in list_extension_api_methods(api_cls):
|
|
||||||
if method.method_id == method_id:
|
|
||||||
return method
|
|
||||||
raise KeyError(f"Unknown extension API method '{method_id}'.")
|
|
||||||
|
|
||||||
|
|
||||||
def extension_api_contract(
|
|
||||||
api_cls: type[ExtensionHostAPI] = ExtensionHostAPI,
|
|
||||||
) -> dict[str, object]:
|
|
||||||
return {
|
|
||||||
"version": 1,
|
|
||||||
"methods": [
|
|
||||||
{
|
|
||||||
"id": method.method_id,
|
|
||||||
"namespace": method.namespace,
|
|
||||||
"name": method.name,
|
|
||||||
"python_name": method.python_name,
|
|
||||||
"host_interface": method.host_interface,
|
|
||||||
"host_name": method.host_name,
|
|
||||||
"sdk_name": method.sdk_name,
|
|
||||||
"sdk_qualified_name": method.sdk_qualified_name,
|
|
||||||
"description": method.description,
|
|
||||||
"required_permission": method.required_permission,
|
|
||||||
"require_auth": method.require_auth,
|
|
||||||
"request_schema": method.request_model.schema(
|
|
||||||
ref_template="#/definitions/{model}"
|
|
||||||
),
|
|
||||||
"response_schema": method.response_model.schema(
|
|
||||||
ref_template="#/definitions/{model}"
|
|
||||||
),
|
|
||||||
}
|
|
||||||
for method in list_extension_api_methods(api_cls)
|
|
||||||
],
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def _get_method_models(
|
|
||||||
function: Callable[..., object],
|
|
||||||
) -> tuple[type[BaseModel], type[BaseModel]]:
|
|
||||||
signature = inspect.signature(function)
|
|
||||||
request_parameters = [
|
|
||||||
parameter
|
|
||||||
for parameter in signature.parameters.values()
|
|
||||||
if parameter.name != "self"
|
|
||||||
]
|
|
||||||
if len(request_parameters) != 1:
|
|
||||||
raise TypeError(
|
|
||||||
f"Extension API method '{function.__name__}' must accept one request model."
|
|
||||||
)
|
)
|
||||||
|
|
||||||
hints = get_type_hints(function)
|
|
||||||
request_model = hints.get(request_parameters[0].name)
|
|
||||||
response_model = hints.get("return")
|
|
||||||
|
|
||||||
if not _is_pydantic_model(request_model):
|
|
||||||
raise TypeError(
|
|
||||||
f"Extension API method '{function.__name__}' request must be a BaseModel."
|
|
||||||
)
|
|
||||||
if not _is_pydantic_model(response_model):
|
|
||||||
raise TypeError(
|
|
||||||
f"Extension API method '{function.__name__}' response must be a BaseModel."
|
|
||||||
)
|
|
||||||
|
|
||||||
return cast(type[BaseModel], request_model), cast(type[BaseModel], response_model)
|
|
||||||
|
|
||||||
|
|
||||||
def _is_pydantic_model(value: object) -> bool:
|
|
||||||
return isinstance(value, type) and issubclass(value, BaseModel)
|
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ from collections.abc import Iterable
|
|||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from lnbits.core.models.extensions import ExtensionPermission, InstallableExtension
|
from lnbits.core.models.extensions import ExtensionPermission, InstallableExtension
|
||||||
from lnbits.core.wasm_ext.api.host import extension_api_permission_ids
|
from lnbits.core.wasm_ext.api.registry import extension_api_permission_ids
|
||||||
|
|
||||||
|
|
||||||
def validate_extension_permissions(
|
def validate_extension_permissions(
|
||||||
|
|||||||
@@ -0,0 +1,199 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import inspect
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from functools import wraps
|
||||||
|
from typing import Any, TypeVar, cast, get_type_hints
|
||||||
|
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
from .models import ExtensionAPIMethod, ExtensionAPIMethodExport
|
||||||
|
|
||||||
|
_EXTENSION_API_METHOD_ATTR = "__lnbits_extension_api_method__"
|
||||||
|
_EXTENSION_RUNTIME_PERMISSION_IDS = {"ui.camera.scan_qr"}
|
||||||
|
_RequestModel = TypeVar("_RequestModel", bound=BaseModel)
|
||||||
|
_ResponseModel = TypeVar("_ResponseModel", bound=BaseModel)
|
||||||
|
|
||||||
|
|
||||||
|
def extension_api_method(
|
||||||
|
*,
|
||||||
|
method_id: str,
|
||||||
|
namespace: str,
|
||||||
|
name: str,
|
||||||
|
host_name: str,
|
||||||
|
sdk_name: str,
|
||||||
|
description: str,
|
||||||
|
host_interface: str = "host",
|
||||||
|
required_permission: str | None = None,
|
||||||
|
require_auth: bool = True,
|
||||||
|
) -> Callable[
|
||||||
|
[Callable[[Any, _RequestModel], Awaitable[_ResponseModel]]],
|
||||||
|
Callable[[Any, _RequestModel], Awaitable[_ResponseModel]],
|
||||||
|
]:
|
||||||
|
export = ExtensionAPIMethodExport(
|
||||||
|
method_id=method_id,
|
||||||
|
namespace=namespace,
|
||||||
|
name=name,
|
||||||
|
host_interface=host_interface,
|
||||||
|
host_name=host_name,
|
||||||
|
sdk_name=sdk_name,
|
||||||
|
description=description,
|
||||||
|
required_permission=required_permission,
|
||||||
|
require_auth=require_auth,
|
||||||
|
)
|
||||||
|
|
||||||
|
def decorator(
|
||||||
|
function: Callable[[Any, _RequestModel], Awaitable[_ResponseModel]],
|
||||||
|
) -> Callable[[Any, _RequestModel], Awaitable[_ResponseModel]]:
|
||||||
|
@wraps(function)
|
||||||
|
async def wrapper(self: Any, request: _RequestModel) -> _ResponseModel:
|
||||||
|
api = getattr(self, "api", self)
|
||||||
|
if require_auth and not api.has_authenticated_context():
|
||||||
|
raise PermissionError(
|
||||||
|
f"Extension API method '{method_id}' requires authentication."
|
||||||
|
)
|
||||||
|
api.require_permission(required_permission)
|
||||||
|
return await function(self, request)
|
||||||
|
|
||||||
|
setattr(wrapper, _EXTENSION_API_METHOD_ATTR, export)
|
||||||
|
return wrapper
|
||||||
|
|
||||||
|
return decorator
|
||||||
|
|
||||||
|
|
||||||
|
def list_extension_api_methods(
|
||||||
|
api_cls: type[Any] | None = None,
|
||||||
|
) -> list[ExtensionAPIMethod]:
|
||||||
|
api_cls = _default_api_cls(api_cls)
|
||||||
|
methods: list[ExtensionAPIMethod] = []
|
||||||
|
|
||||||
|
for prefix, method_cls in _extension_api_method_sources(api_cls):
|
||||||
|
for python_name, function in inspect.getmembers(method_cls, inspect.isfunction):
|
||||||
|
export = getattr(function, _EXTENSION_API_METHOD_ATTR, None)
|
||||||
|
if not export:
|
||||||
|
continue
|
||||||
|
|
||||||
|
request_model, response_model = _get_method_models(function)
|
||||||
|
methods.append(
|
||||||
|
ExtensionAPIMethod(
|
||||||
|
method_id=export.method_id,
|
||||||
|
namespace=export.namespace,
|
||||||
|
name=export.name,
|
||||||
|
python_name=f"{prefix}.{python_name}" if prefix else python_name,
|
||||||
|
host_interface=export.host_interface,
|
||||||
|
host_name=export.host_name,
|
||||||
|
sdk_name=export.sdk_name,
|
||||||
|
description=export.description,
|
||||||
|
request_model=request_model,
|
||||||
|
response_model=response_model,
|
||||||
|
required_permission=export.required_permission,
|
||||||
|
require_auth=export.require_auth,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return sorted(methods, key=lambda method: method.method_id)
|
||||||
|
|
||||||
|
|
||||||
|
def extension_api_permission_ids(api_cls: type[Any] | None = None) -> set[str]:
|
||||||
|
permissions = {
|
||||||
|
method.required_permission
|
||||||
|
for method in list_extension_api_methods(api_cls)
|
||||||
|
if method.required_permission
|
||||||
|
}
|
||||||
|
permissions.update(_EXTENSION_RUNTIME_PERMISSION_IDS)
|
||||||
|
return permissions
|
||||||
|
|
||||||
|
|
||||||
|
def get_extension_api_method(
|
||||||
|
method_id: str,
|
||||||
|
api_cls: type[Any] | None = None,
|
||||||
|
) -> ExtensionAPIMethod:
|
||||||
|
for method in list_extension_api_methods(api_cls):
|
||||||
|
if method.method_id == method_id:
|
||||||
|
return method
|
||||||
|
raise KeyError(f"Unknown extension API method '{method_id}'.")
|
||||||
|
|
||||||
|
|
||||||
|
def extension_api_contract(api_cls: type[Any] | None = None) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"version": 1,
|
||||||
|
"methods": [
|
||||||
|
{
|
||||||
|
"id": method.method_id,
|
||||||
|
"namespace": method.namespace,
|
||||||
|
"name": method.name,
|
||||||
|
"python_name": method.python_name,
|
||||||
|
"host_interface": method.host_interface,
|
||||||
|
"host_name": method.host_name,
|
||||||
|
"sdk_name": method.sdk_name,
|
||||||
|
"sdk_qualified_name": method.sdk_qualified_name,
|
||||||
|
"description": method.description,
|
||||||
|
"required_permission": method.required_permission,
|
||||||
|
"require_auth": method.require_auth,
|
||||||
|
"request_schema": method.request_model.schema(
|
||||||
|
ref_template="#/definitions/{model}"
|
||||||
|
),
|
||||||
|
"response_schema": method.response_model.schema(
|
||||||
|
ref_template="#/definitions/{model}"
|
||||||
|
),
|
||||||
|
}
|
||||||
|
for method in list_extension_api_methods(api_cls)
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _default_api_cls(api_cls: type[Any] | None) -> type[Any]:
|
||||||
|
if api_cls is not None:
|
||||||
|
return api_cls
|
||||||
|
|
||||||
|
from .host import ExtensionHostAPI
|
||||||
|
|
||||||
|
return ExtensionHostAPI
|
||||||
|
|
||||||
|
|
||||||
|
def _extension_api_method_sources(
|
||||||
|
api_cls: type[Any],
|
||||||
|
) -> list[tuple[str, type[Any]]]:
|
||||||
|
sources: list[tuple[str, type[Any]]] = [("", api_cls)]
|
||||||
|
|
||||||
|
from .host import ExtensionHostAPI
|
||||||
|
|
||||||
|
if issubclass(api_cls, ExtensionHostAPI):
|
||||||
|
from .utils import extension_api_utils_method_classes
|
||||||
|
|
||||||
|
sources.extend(extension_api_utils_method_classes().items())
|
||||||
|
return sources
|
||||||
|
|
||||||
|
|
||||||
|
def _get_method_models(
|
||||||
|
function: Callable[..., object],
|
||||||
|
) -> tuple[type[BaseModel], type[BaseModel]]:
|
||||||
|
signature = inspect.signature(function)
|
||||||
|
request_parameters = [
|
||||||
|
parameter
|
||||||
|
for parameter in signature.parameters.values()
|
||||||
|
if parameter.name != "self"
|
||||||
|
]
|
||||||
|
if len(request_parameters) != 1:
|
||||||
|
raise TypeError(
|
||||||
|
f"Extension API method '{function.__name__}' must accept one request model."
|
||||||
|
)
|
||||||
|
|
||||||
|
hints = get_type_hints(function)
|
||||||
|
request_model = hints.get(request_parameters[0].name)
|
||||||
|
response_model = hints.get("return")
|
||||||
|
|
||||||
|
if not _is_pydantic_model(request_model):
|
||||||
|
raise TypeError(
|
||||||
|
f"Extension API method '{function.__name__}' request must be a BaseModel."
|
||||||
|
)
|
||||||
|
if not _is_pydantic_model(response_model):
|
||||||
|
raise TypeError(
|
||||||
|
f"Extension API method '{function.__name__}' response must be a BaseModel."
|
||||||
|
)
|
||||||
|
|
||||||
|
return cast(type[BaseModel], request_model), cast(type[BaseModel], response_model)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_pydantic_model(value: object) -> bool:
|
||||||
|
return isinstance(value, type) and issubclass(value, BaseModel)
|
||||||
@@ -7,7 +7,9 @@ from typing import Any
|
|||||||
|
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
from .host import ExtensionAPIMethod, ExtensionHostAPI, list_extension_api_methods
|
from .host import ExtensionHostAPI
|
||||||
|
from .models import ExtensionAPIMethod
|
||||||
|
from .registry import list_extension_api_methods
|
||||||
|
|
||||||
HostImport = Callable[..., Awaitable[dict[str, Any]]]
|
HostImport = Callable[..., Awaitable[dict[str, Any]]]
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ import time
|
|||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
from .host import extension_api_method
|
|
||||||
from .models import (
|
from .models import (
|
||||||
Bolt11Request,
|
Bolt11Request,
|
||||||
CurrencyConvertRequest,
|
CurrencyConvertRequest,
|
||||||
@@ -29,6 +28,7 @@ from .models import (
|
|||||||
VerifyPreimageRequest,
|
VerifyPreimageRequest,
|
||||||
VerifyPreimageResponse,
|
VerifyPreimageResponse,
|
||||||
)
|
)
|
||||||
|
from .registry import extension_api_method
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from .host import ExtensionHostAPI
|
from .host import ExtensionHostAPI
|
||||||
|
|||||||
@@ -5,8 +5,8 @@ import re
|
|||||||
from collections.abc import Mapping
|
from collections.abc import Mapping
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from ..api.host import list_extension_api_methods
|
|
||||||
from ..api.models import EmptyRequest
|
from ..api.models import EmptyRequest
|
||||||
|
from ..api.registry import list_extension_api_methods
|
||||||
from ..api.runtime import ExtensionAPIHost
|
from ..api.runtime import ExtensionAPIHost
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user