diff --git a/src/plyngent/agent/tools.py b/src/plyngent/agent/tools.py new file mode 100644 index 0000000..f7070ef --- /dev/null +++ b/src/plyngent/agent/tools.py @@ -0,0 +1,201 @@ +from __future__ import annotations + +import inspect +import types +from collections.abc import Awaitable, Callable, Mapping +from typing import Any, cast, get_args, get_origin, get_type_hints, overload + +import msgspec + +from plyngent.lmproto.openai_compatible.model import ToolFunction, ToolFunctionItem +from plyngent.typedef import JSONSchema # noqa: TC001 + +type ToolHandler = Callable[..., Any | Awaitable[Any]] + +_PRIMITIVE_SCHEMA: dict[type, JSONSchema] = { + str: {"type": "string"}, + int: {"type": "integer"}, + float: {"type": "number"}, + bool: {"type": "boolean"}, +} + + +class ToolDefinition: + """A registered tool: schema for the model plus a callable handler.""" + + name: str + description: str + parameters: JSONSchema + handler: ToolHandler + + def __init__( + self, + name: str, + description: str, + parameters: JSONSchema, + handler: ToolHandler, + ) -> None: + self.name = name + self.description = description + self.parameters = parameters + self.handler = handler + + def to_tool_item(self) -> ToolFunctionItem: + return ToolFunctionItem( + function=ToolFunction( + name=self.name, + description=self.description or msgspec.UNSET, + parameters=self.parameters or msgspec.UNSET, + ) + ) + + +def _annotation_to_schema(annotation: object) -> JSONSchema: + if annotation is inspect.Parameter.empty or annotation is Any: + return {} + origin = get_origin(annotation) + if origin is types.UnionType: + args = [a for a in get_args(annotation) if a is not type(None)] + return _annotation_to_schema(args[0]) if len(args) == 1 else {} + if origin is list: + args = get_args(annotation) + items = _annotation_to_schema(args[0]) if args else {} + return {"type": "array", "items": items or True} + if origin is dict: + return {"type": "object"} + if isinstance(annotation, type) and annotation in _PRIMITIVE_SCHEMA: + return dict(_PRIMITIVE_SCHEMA[annotation]) + return {} + + +def schema_from_callable(func: ToolHandler) -> JSONSchema: + """Build a JSON Schema object for ``func`` parameters from type hints.""" + hints = get_type_hints(func) + signature = inspect.signature(func) + properties: dict[str, JSONSchema] = {} + required: list[str] = [] + for name, param in signature.parameters.items(): + if name in {"self", "cls"}: + continue + if param.kind in { + inspect.Parameter.VAR_POSITIONAL, + inspect.Parameter.VAR_KEYWORD, + }: + continue + annotation = hints.get(name, param.annotation) + properties[name] = _annotation_to_schema(annotation) + if param.default is inspect.Parameter.empty: + required.append(name) + schema: JSONSchema = {"type": "object", "properties": properties} + if required: + schema["required"] = required + return schema + + +def _build_definition( + func: ToolHandler, + *, + name: str | None, + description: str | None, +) -> ToolDefinition: + tool_name = name or func.__name__ + tool_description = description if description is not None else (inspect.getdoc(func) or "") + return ToolDefinition( + name=tool_name, + description=tool_description, + parameters=schema_from_callable(func), + handler=func, + ) + + +@overload +def tool[**PS, R](func: Callable[PS, R], /) -> ToolDefinition: ... + + +@overload +def tool[**PS, R]( + func: None = None, + /, + *, + name: str | None = None, + description: str | None = None, +) -> Callable[[Callable[PS, R]], ToolDefinition]: ... + + +def tool[**PS, R]( + func: Callable[PS, R] | None = None, + /, + *, + name: str | None = None, + description: str | None = None, +) -> ToolDefinition | Callable[[Callable[PS, R]], ToolDefinition]: + """Register a function as an agent tool (decorator). + + Schema is inferred from type hints; description defaults to the docstring. + """ + + def decorator(fn: Callable[PS, R]) -> ToolDefinition: + return _build_definition(fn, name=name, description=description) + + if func is not None: + return decorator(func) + return decorator + + +class ToolRegistry: + """Name → tool definition map with execution helpers.""" + + _tools: dict[str, ToolDefinition] + + def __init__(self, tools: Mapping[str, ToolDefinition] | list[ToolDefinition] | None = None) -> None: + self._tools = {} + if tools is None: + return + if isinstance(tools, list): + for item in tools: + self.register(item) + else: + for item in tools.values(): + self.register(item) + + def register(self, definition: ToolDefinition) -> None: + self._tools[definition.name] = definition + + def get(self, name: str) -> ToolDefinition | None: + return self._tools.get(name) + + def tool_items(self) -> list[ToolFunctionItem]: + return [t.to_tool_item() for t in self._tools.values()] + + def __contains__(self, name: str) -> bool: + return name in self._tools + + def __len__(self) -> int: + return len(self._tools) + + async def _invoke(self, definition: ToolDefinition, args: dict[str, object]) -> str: + try: + result = definition.handler(**args) + if inspect.isawaitable(result): + result = await result + except TypeError as exc: + return f"error: invalid tool arguments: {exc}" + except Exception as exc: # noqa: BLE001 — surface tool failures to the model + return f"error: tool {definition.name!r} failed: {exc}" + if isinstance(result, str): + return result + return msgspec.json.encode(result).decode() + + async def execute(self, name: str, arguments_json: str) -> str: + """Run a tool by name; returns a string result (errors become error text).""" + definition = self._tools.get(name) + if definition is None: + return f"error: unknown tool {name!r}" + try: + raw_args: object = msgspec.json.decode(arguments_json.encode()) + except (msgspec.DecodeError, UnicodeEncodeError) as exc: + return f"error: invalid tool arguments JSON: {exc}" + if not isinstance(raw_args, dict): + return "error: tool arguments must be a JSON object" + args = {str(key): value for key, value in cast("dict[object, object]", raw_args).items()} + return await self._invoke(definition, args)