_llm called env_manager.set_env() at import time, which connected to a private ArangoDB to fetch secrets. That single line meant a fresh clone could not start, regardless of what else was configured. Both packages also lived in separate private repos and were gitignored here, so the code shipped without them. packages/llm/ is 780 lines against _llm's 1750. Dropped as unused by this project (measured, zero call sites): token counting and message trimming, image/vision handling, make_summary, the ollama-specific paths, the query/user_input/context argument style, and the self-mutating provider_quirks.json cache. Kept and reworked: - tools.py, the docstring -> JSON-schema tool registry, which has no equivalent in the cuj-fup client and which llm_tools.py depends on entirely. - The provider quirks that actually matter: vLLM-only extra_body fields stripped for hosted providers, enable_thinking disabled at template level when think is off, reasoning models (o1/o3/o4/gpt-5) switched to max_completion_tokens. Adopted from cuj-fup's client: LLMConfig as a dataclass instead of 20 constructor kwargs, the SDK's native max_retries instead of hand-rolled backoff, and error messages that name the likely cause. Fixes a latent bug: Optional[list[str]] parameters were advertised to the model as strings, because get_origin(Optional[X]) is Union, so neither the schema mapping nor the list coercion in execute_tool fired. `parties`, `people` and `focus_ids` were all affected. Also drops a dead SELECT-only guard in execute_tool that keyed on a parameter name (`sql_query`) that no tool has ever used. Real SQL hardening is tracked separately. Verified: `import backend.app` succeeds with all external network blocked and zero outbound connection attempts; all 12 tools register; live vLLM calls confirmed for plain generation, structured output via format=, tool execution, and the error-returns-a-string contract that call sites branch on. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>main
parent
09c1b3d79c
commit
843db25211
16 changed files with 885 additions and 16 deletions
@ -0,0 +1,64 @@ |
||||
"""Coloured console output for the ingestion and chat pipelines. |
||||
|
||||
Replaces the small private `colorprinter` package the predecessor depended on. |
||||
Colour is suppressed automatically when stdout is not a terminal, and when the |
||||
NO_COLOR convention (https://no-color.org) is in effect — so piping output to a |
||||
file or a systemd journal yields clean text. |
||||
""" |
||||
from __future__ import annotations |
||||
|
||||
import os |
||||
import sys |
||||
from typing import Any |
||||
|
||||
_CODES = { |
||||
"red": "\033[31m", |
||||
"green": "\033[32m", |
||||
"yellow": "\033[33m", |
||||
"blue": "\033[34m", |
||||
"purple": "\033[35m", |
||||
} |
||||
_RESET = "\033[0m" |
||||
|
||||
|
||||
def _colour_enabled() -> bool: |
||||
if os.environ.get("NO_COLOR") is not None: |
||||
return False |
||||
if os.environ.get("FORCE_COLOR"): |
||||
return True |
||||
return sys.stdout.isatty() |
||||
|
||||
|
||||
def _emit(colour: str, *args: Any, **kwargs: Any) -> None: |
||||
text = " ".join(str(a) for a in args) |
||||
if _colour_enabled(): |
||||
text = f"{_CODES[colour]}{text}{_RESET}" |
||||
print(text, **kwargs) |
||||
|
||||
|
||||
def print_red(*args: Any, **kwargs: Any) -> None: |
||||
"""Errors and blocked operations.""" |
||||
_emit("red", *args, **kwargs) |
||||
|
||||
|
||||
def print_yellow(*args: Any, **kwargs: Any) -> None: |
||||
"""Warnings, timings, and retries.""" |
||||
_emit("yellow", *args, **kwargs) |
||||
|
||||
|
||||
def print_green(*args: Any, **kwargs: Any) -> None: |
||||
"""Successful completion of a step.""" |
||||
_emit("green", *args, **kwargs) |
||||
|
||||
|
||||
def print_blue(*args: Any, **kwargs: Any) -> None: |
||||
"""Informational trace, e.g. the SQL a tool is about to run.""" |
||||
_emit("blue", *args, **kwargs) |
||||
|
||||
|
||||
def print_purple(*args: Any, **kwargs: Any) -> None: |
||||
"""Secondary trace.""" |
||||
_emit("purple", *args, **kwargs) |
||||
|
||||
|
||||
__all__ = ["print_red", "print_yellow", "print_green", "print_blue", "print_purple"] |
||||
@ -0,0 +1,39 @@ |
||||
"""Provider-agnostic LLM client and tool registry. |
||||
|
||||
Replaces the `_llm` package the predecessor project depended on, which reached out |
||||
to a private server at import time and so could not be run by anyone else. |
||||
|
||||
from packages.llm import LLM, register_tool, get_tools |
||||
|
||||
@register_tool |
||||
def search(query: str) -> str: |
||||
'''Search the corpus. |
||||
|
||||
Args: |
||||
query: What to look for. |
||||
''' |
||||
... |
||||
|
||||
llm = LLM(base_url=..., model=..., tools=get_tools()) |
||||
reply = llm.generate(messages=[{"role": "user", "content": "..."}]) |
||||
""" |
||||
from .client import LLM, ChatCompletionMessage |
||||
from .config import LLMConfig |
||||
from .tools import ( |
||||
TOOL_REGISTRY, |
||||
execute_tool, |
||||
get_tools, |
||||
parse_function_call_arguments, |
||||
register_tool, |
||||
) |
||||
|
||||
__all__ = [ |
||||
"LLM", |
||||
"LLMConfig", |
||||
"ChatCompletionMessage", |
||||
"register_tool", |
||||
"get_tools", |
||||
"execute_tool", |
||||
"parse_function_call_arguments", |
||||
"TOOL_REGISTRY", |
||||
] |
||||
@ -0,0 +1,390 @@ |
||||
"""A thin, provider-agnostic wrapper over the OpenAI-compatible chat API. |
||||
|
||||
Works against self-hosted vLLM, OpenAI, OpenRouter, Berget, Gemini's compatibility |
||||
endpoint — anything speaking the OpenAI chat protocol. The differences between them |
||||
that actually bite in production are handled in one place here: |
||||
|
||||
* vLLM accepts ``extra_body`` sampler knobs (``repetition_penalty``, |
||||
``chat_template_kwargs``) that hosted providers reject with a 4xx. |
||||
* OpenAI reasoning models (o1/o3/o4/gpt-5) require ``max_completion_tokens`` |
||||
instead of ``max_tokens``, and refuse any temperature other than 1. |
||||
* Some models emit ``<think>`` blocks inline in ``content`` rather than in |
||||
``reasoning_content``; those must never reach a user. |
||||
|
||||
Error contract, preserved from the predecessor: on API failure ``generate`` |
||||
returns an error *string* rather than raising. Call sites branch on |
||||
``isinstance(response, str)``. |
||||
""" |
||||
from __future__ import annotations |
||||
|
||||
import json |
||||
import os |
||||
import re |
||||
import traceback |
||||
from typing import Any, Dict, Generator, List, Optional, Tuple, Type |
||||
|
||||
from openai import OpenAI |
||||
from openai.types.chat import ChatCompletion |
||||
from openai.types.chat.chat_completion import Choice |
||||
from openai.types.chat.chat_completion_message import ( |
||||
ChatCompletionMessage as _OpenAIChatCompletionMessage, |
||||
) |
||||
from pydantic import BaseModel |
||||
|
||||
from .tools import execute_tool, parse_function_call_arguments |
||||
|
||||
_THINK_RE = re.compile(r"<think>.*?</think>\s*", re.DOTALL | re.IGNORECASE) |
||||
|
||||
# Models that reject `max_tokens` and any temperature but 1. |
||||
_REASONING_PREFIXES = ("o1", "o3", "o4", "gpt-5") |
||||
|
||||
# Sampler fields only self-hosted vLLM understands. |
||||
_VLLM_ONLY_EXTRA_BODY = ("repetition_penalty", "chat_template_kwargs") |
||||
|
||||
|
||||
class ChatCompletionMessage(_OpenAIChatCompletionMessage): |
||||
"""Assistant message, extended with the structured-output fields. |
||||
|
||||
When ``generate(format=...)`` is used, ``content`` holds the *parsed model |
||||
instance* rather than text, and ``parsed`` / ``content_text`` carry the |
||||
instance and its raw JSON respectively. |
||||
""" |
||||
|
||||
model_config = {"extra": "allow"} |
||||
|
||||
|
||||
def _strip_think(text: str) -> str: |
||||
return _THINK_RE.sub("", text).strip() |
||||
|
||||
|
||||
def _is_reasoning_model(model: str) -> bool: |
||||
name = (model or "").split("/")[-1].lower() |
||||
return name.startswith(_REASONING_PREFIXES) |
||||
|
||||
|
||||
class LLM: |
||||
"""One configured connection to a chat model. |
||||
|
||||
Args: |
||||
system_message: Seeds ``self.messages`` when no explicit history is passed. |
||||
temperature: Default sampling temperature; overridable per call. |
||||
model: Model identifier sent to the provider. |
||||
base_url: OpenAI-compatible endpoint, including the ``/v1`` suffix. |
||||
api_key: Provider key. Its presence also marks this as an *external* |
||||
provider, which suppresses the vLLM-only ``extra_body`` fields. |
||||
think: Default reasoning mode. When false on self-hosted vLLM, chain-of-thought |
||||
is disabled at the template level so no reasoning tokens are generated. |
||||
""" |
||||
|
||||
def __init__( |
||||
self, |
||||
system_message: str = "You are an assistant.", |
||||
temperature: float = 0.01, |
||||
model: Optional[str] = None, |
||||
max_length_answer: int = 3000, |
||||
messages: Optional[List[dict]] = None, |
||||
chat: bool = True, |
||||
tools: Optional[list] = None, |
||||
think: bool = False, |
||||
timeout: int = 240, |
||||
silent: bool = False, |
||||
presence_penalty: float = 0.3, |
||||
top_p: float = 0.9, |
||||
extra_body: Optional[Dict[str, Any]] = None, |
||||
base_url: Optional[str] = None, |
||||
api_key: Optional[str] = None, |
||||
max_retries: int = 4, |
||||
) -> None: |
||||
self.model = model or os.getenv("LLM_MODEL", "smart") |
||||
self.system_message = system_message |
||||
self.messages = messages or [{"role": "system", "content": system_message}] |
||||
self.max_length_answer = max_length_answer |
||||
self.chat = chat |
||||
self.tools = tools or [] |
||||
self.think = think |
||||
self.silent = silent |
||||
self.options = { |
||||
"temperature": temperature, |
||||
"presence_penalty": presence_penalty, |
||||
"top_p": top_p, |
||||
} |
||||
# repetition_penalty > 1.0 damps already-seen tokens; 1.2 breaks generation |
||||
# loops without measurably hurting quality. |
||||
self.extra_body = extra_body if extra_body is not None else {"repetition_penalty": 1.2} |
||||
|
||||
self._api_key = api_key |
||||
self.base_url = base_url or os.getenv("LLM_DIRECT_URL") or "" |
||||
if not self.base_url: |
||||
raise ValueError( |
||||
"No LLM endpoint configured. Pass base_url= or set LLM_DIRECT_URL." |
||||
) |
||||
|
||||
# max_retries covers transient 408/409/429/5xx with exponential backoff. |
||||
# Above the SDK default of 2 because a self-hosted vLLM engine crash returns |
||||
# 500 and usually restarts within seconds — the extra attempts ride it out. |
||||
self.client = OpenAI( |
||||
base_url=self.base_url, |
||||
api_key=self._api_key or os.getenv("LLM_BEARER", "NONE"), |
||||
timeout=timeout, |
||||
max_retries=max_retries, |
||||
) |
||||
|
||||
# -- request assembly ----------------------------------------------------- |
||||
|
||||
@property |
||||
def _is_external_provider(self) -> bool: |
||||
"""A caller-supplied key means a hosted provider, not our own vLLM.""" |
||||
return bool(self._api_key) |
||||
|
||||
def _build_extra_body(self, think: Optional[bool]) -> Optional[dict]: |
||||
body = dict(self.extra_body or {}) |
||||
|
||||
effective_think = self.think if think is None else think |
||||
if not effective_think and not self._is_external_provider: |
||||
ctk = dict(body.get("chat_template_kwargs") or {}) |
||||
ctk["enable_thinking"] = False |
||||
body["chat_template_kwargs"] = ctk |
||||
|
||||
if self._is_external_provider: |
||||
for key in _VLLM_ONLY_EXTRA_BODY: |
||||
body.pop(key, None) |
||||
|
||||
return body or None |
||||
|
||||
def _sampling_kwargs(self, model: str, temperature: Optional[float], |
||||
max_tokens: Optional[int]) -> dict: |
||||
temp = self.options["temperature"] if temperature is None else temperature |
||||
limit = max_tokens or self.max_length_answer |
||||
|
||||
if _is_reasoning_model(model): |
||||
# These models accept only the default temperature and a different token field. |
||||
return {"max_completion_tokens": limit} |
||||
|
||||
return { |
||||
"temperature": temp, |
||||
"top_p": self.options["top_p"], |
||||
"max_tokens": limit, |
||||
} |
||||
|
||||
def _create(self, **kwargs) -> ChatCompletion: |
||||
"""Call the API, retrying once if the token-limit field name is rejected.""" |
||||
try: |
||||
return self.client.chat.completions.create(**kwargs) |
||||
except Exception as exc: |
||||
swapped = _swap_token_param(kwargs, exc) |
||||
if swapped is None: |
||||
raise _friendly_error(exc, self.base_url) from exc |
||||
return self.client.chat.completions.create(**swapped) |
||||
|
||||
# -- public API ----------------------------------------------------------- |
||||
|
||||
def generate( |
||||
self, |
||||
query: Optional[str] = None, |
||||
*, |
||||
messages: Optional[List[dict]] = None, |
||||
tools: Optional[list] = None, |
||||
model: Optional[str] = None, |
||||
temperature: Optional[float] = None, |
||||
format: Optional[Type[BaseModel]] = None, |
||||
stream: bool = False, |
||||
think: Optional[bool] = None, |
||||
max_tokens: Optional[int] = None, |
||||
extra_body: Optional[Dict[str, Any]] = None, |
||||
auto_execute_tools: bool = True, |
||||
): |
||||
"""Run one completion. |
||||
|
||||
Returns a :class:`ChatCompletionMessage`, a generator of ``(kind, chunk)`` |
||||
pairs when ``stream`` is set, or an error string when the call fails. |
||||
|
||||
Args: |
||||
query: Convenience for a single user turn; ignored if ``messages`` is given. |
||||
messages: Full conversation to send. Also becomes this instance's history. |
||||
tools: Tool schemas to advertise. Defaults to the instance's tools. |
||||
format: Pydantic model requesting structured output via JSON schema. |
||||
auto_execute_tools: Execute returned tool calls and append their results |
||||
to the history. It does *not* continue the conversation — the caller |
||||
decides whether to call again. |
||||
""" |
||||
if messages is not None: |
||||
self.messages = list(messages) |
||||
elif query is not None: |
||||
self.messages.append({"role": "user", "content": query}) |
||||
|
||||
resolved_model = model or self.model |
||||
if extra_body is not None: |
||||
self.extra_body = extra_body |
||||
|
||||
try: |
||||
if format is not None: |
||||
return self._generate_structured( |
||||
resolved_model, format, temperature, max_tokens, think |
||||
) |
||||
|
||||
request = { |
||||
"model": resolved_model, |
||||
"messages": self.messages, |
||||
"extra_body": self._build_extra_body(think), |
||||
**self._sampling_kwargs(resolved_model, temperature, max_tokens), |
||||
} |
||||
tools_to_use = self.tools if tools is None else tools |
||||
if tools_to_use: |
||||
request["tools"] = tools_to_use |
||||
if stream: |
||||
request["stream"] = True |
||||
return self._read_stream(self.client.chat.completions.create(**request)) |
||||
|
||||
message = self._create(**request).choices[0].message |
||||
|
||||
if auto_execute_tools and getattr(message, "tool_calls", None): |
||||
self._run_tool_calls(message.tool_calls) |
||||
|
||||
if isinstance(message.content, str): |
||||
message.content = _strip_think(message.content) |
||||
|
||||
self.messages.append({"role": "assistant", "content": message.content}) |
||||
if not self.chat: |
||||
self.messages = self.messages[:1] |
||||
return message |
||||
|
||||
except Exception as exc: |
||||
if not self.silent: |
||||
traceback.print_exc() |
||||
return f"LLM request failed: {exc}" |
||||
|
||||
# -- structured output ---------------------------------------------------- |
||||
|
||||
def _generate_structured( |
||||
self, |
||||
model: str, |
||||
format: Type[BaseModel], |
||||
temperature: Optional[float], |
||||
max_tokens: Optional[int], |
||||
think: Optional[bool], |
||||
) -> ChatCompletionMessage: |
||||
# vLLM's json_schema mode rejects `role: tool` turns, so fold them into |
||||
# user turns. Done on a copy — the caller's history stays intact. |
||||
messages = [ |
||||
{"role": "user", "content": f"Tool output:\n{m.get('content', '')}"} |
||||
if m.get("role") == "tool" else m |
||||
for m in self.messages |
||||
] |
||||
|
||||
response = self._create( |
||||
model=model, |
||||
messages=messages, |
||||
extra_body=self._build_extra_body(think), |
||||
response_format={ |
||||
"type": "json_schema", |
||||
"json_schema": {"name": format.__name__, "schema": format.model_json_schema()}, |
||||
}, |
||||
**self._sampling_kwargs(model, temperature, max_tokens), |
||||
) |
||||
|
||||
content_text = response.choices[0].message.content or "" |
||||
parsed = format.model_validate_json(_extract_json(content_text)) |
||||
|
||||
message = ChatCompletionMessage.model_construct(role="assistant", content=parsed) |
||||
message.parsed = parsed |
||||
message.parsed_dict = parsed.model_dump() |
||||
message.content_text = content_text |
||||
return message |
||||
|
||||
# -- tool execution ------------------------------------------------------- |
||||
|
||||
def _run_tool_calls(self, tool_calls) -> None: |
||||
"""Execute each returned tool call, appending results as `tool` messages. |
||||
|
||||
A failing tool appends its error rather than raising, so the model can see |
||||
what went wrong and correct itself on the next turn. |
||||
""" |
||||
for call in tool_calls: |
||||
fn = getattr(call, "function", None) |
||||
if fn is None: |
||||
continue |
||||
name = getattr(fn, "name", None) |
||||
try: |
||||
args = parse_function_call_arguments(getattr(fn, "arguments", None)) |
||||
result = execute_tool(name, args) |
||||
content = result if isinstance(result, str) else json.dumps(result, ensure_ascii=False) |
||||
except Exception as exc: |
||||
if not self.silent: |
||||
print(f"[llm] tool {name} failed: {exc}") |
||||
content = json.dumps({"error": str(exc)}, ensure_ascii=False) |
||||
self.messages.append({"role": "tool", "name": name or "unknown", "content": content}) |
||||
|
||||
# -- streaming ------------------------------------------------------------ |
||||
|
||||
def _read_stream(self, response) -> Generator[Tuple[str, str], None, None]: |
||||
"""Yield ``("thinking" | "content", text)`` pairs as they arrive. |
||||
|
||||
Reasoning arrives either in ``reasoning_content`` or inline as ``<think>`` |
||||
blocks; both are surfaced as "thinking" so callers can render or drop them. |
||||
""" |
||||
in_think_block = False |
||||
for chunk in response: |
||||
if not chunk.choices: |
||||
continue |
||||
delta = chunk.choices[0].delta |
||||
|
||||
reasoning = getattr(delta, "reasoning_content", None) |
||||
if reasoning: |
||||
yield "thinking", reasoning |
||||
|
||||
text = getattr(delta, "content", None) |
||||
if not text: |
||||
continue |
||||
|
||||
if "<think>" in text: |
||||
in_think_block = True |
||||
text = text.split("<think>", 1)[0] |
||||
if "</think>" in text: |
||||
in_think_block = False |
||||
text = text.split("</think>", 1)[1] |
||||
|
||||
if not text: |
||||
continue |
||||
yield ("thinking" if in_think_block else "content"), text |
||||
|
||||
|
||||
# -- module helpers ----------------------------------------------------------- |
||||
|
||||
|
||||
def _extract_json(text: str) -> str: |
||||
"""Pull the JSON object out of a response that may carry stray prose.""" |
||||
text = _strip_think(text).strip() |
||||
if text.startswith("{"): |
||||
return text |
||||
match = re.search(r"\{.*\}", text, re.DOTALL) |
||||
return match.group(0) if match else text |
||||
|
||||
|
||||
def _swap_token_param(kwargs: dict, exc: Exception) -> Optional[dict]: |
||||
"""Retry payload with the other token-limit field, if that's what was rejected. |
||||
|
||||
Providers disagree about `max_tokens` vs `max_completion_tokens`, and the model |
||||
lists that require each keep changing. Rather than track them, react to the error. |
||||
""" |
||||
text = str(exc).lower() |
||||
if "max_tokens" not in text and "max_completion_tokens" not in text: |
||||
return None |
||||
|
||||
alternates = {"max_tokens": "max_completion_tokens", "max_completion_tokens": "max_tokens"} |
||||
for current, replacement in alternates.items(): |
||||
if current in kwargs: |
||||
swapped = dict(kwargs) |
||||
swapped[replacement] = swapped.pop(current) |
||||
swapped.pop("temperature", None) # reasoning models reject it too |
||||
return swapped |
||||
return None |
||||
|
||||
|
||||
def _friendly_error(exc: Exception, base_url: str) -> Exception: |
||||
"""Turn provider errors into something a deployer can act on.""" |
||||
text = str(exc).lower() |
||||
if any(s in text for s in ("connection", "timeout", "refused", "unreachable")): |
||||
return RuntimeError(f"LLM endpoint unreachable at {base_url}. Is the server running? ({exc})") |
||||
if any(s in text for s in ("401", "403", "invalid api key", "unauthorized")): |
||||
return RuntimeError(f"LLM provider rejected the API key. ({exc})") |
||||
return exc |
||||
@ -0,0 +1,90 @@ |
||||
"""Declarative description of an LLM endpoint. |
||||
|
||||
Separating *which endpoint* from *how to call it* is what lets one process talk to |
||||
several providers at once — a self-hosted vLLM for bulk summarisation, a hosted |
||||
model for user-facing chat, and a user's own key for either. |
||||
|
||||
Keys carried here are ephemeral. A user-supplied key arrives on a request, is used |
||||
for that call, and is never written to disk or the database. |
||||
""" |
||||
from __future__ import annotations |
||||
|
||||
import os |
||||
from dataclasses import dataclass, replace |
||||
from typing import Optional |
||||
|
||||
|
||||
@dataclass(frozen=True) |
||||
class LLMConfig: |
||||
"""Everything needed to open a connection to one model. |
||||
|
||||
Args: |
||||
base_url: OpenAI-compatible endpoint including the ``/v1`` suffix. |
||||
model: Default model identifier. |
||||
api_key: Provider key. Empty means self-hosted; its presence also |
||||
suppresses the vLLM-only sampler fields (see ``client``). |
||||
provider: Free-form label used for logging and provider-specific quirks. |
||||
model_fast: Cheaper model for mechanical work (summarising, tagging). |
||||
Falls back to ``model``. |
||||
model_smart: Stronger model for user-facing reasoning. Falls back to ``model``. |
||||
""" |
||||
|
||||
base_url: str |
||||
model: str |
||||
api_key: str = "" |
||||
provider: str = "" |
||||
model_fast: str = "" |
||||
model_smart: str = "" |
||||
temperature: float = 0.01 |
||||
timeout: int = 240 |
||||
max_retries: int = 4 |
||||
|
||||
@classmethod |
||||
def from_env(cls) -> "LLMConfig": |
||||
"""Build the server's default endpoint from environment variables.""" |
||||
base_url = os.getenv("LLM_DIRECT_URL", "") |
||||
if not base_url: |
||||
raise ValueError("LLM_DIRECT_URL is not set; see .env.example") |
||||
return cls( |
||||
base_url=base_url, |
||||
model=os.getenv("LLM_MODEL", "smart"), |
||||
api_key=os.getenv("LLM_BEARER", ""), |
||||
provider="vllm", |
||||
model_fast=os.getenv("LLM_MODEL_FAST", ""), |
||||
model_smart=os.getenv("LLM_MODEL_SMART", ""), |
||||
) |
||||
|
||||
def with_override( |
||||
self, |
||||
*, |
||||
base_url: Optional[str] = None, |
||||
api_key: Optional[str] = None, |
||||
model: Optional[str] = None, |
||||
provider: Optional[str] = None, |
||||
) -> "LLMConfig": |
||||
"""Return a copy pointed at a different provider, for a single request.""" |
||||
changes = { |
||||
k: v |
||||
for k, v in { |
||||
"base_url": base_url, |
||||
"api_key": api_key, |
||||
"model": model, |
||||
"provider": provider, |
||||
}.items() |
||||
if v |
||||
} |
||||
return replace(self, **changes) |
||||
|
||||
def resolve(self, role: str = "default") -> str: |
||||
"""Pick the model for a role, falling back to the default model.""" |
||||
return { |
||||
"fast": self.model_fast, |
||||
"smart": self.model_smart, |
||||
}.get(role, "") or self.model |
||||
|
||||
def __repr__(self) -> str: # never let a key reach a log line |
||||
redacted = "set" if self.api_key else "unset" |
||||
return ( |
||||
f"LLMConfig(provider={self.provider!r}, base_url={self.base_url!r}, " |
||||
f"model={self.model!r}, api_key=<{redacted}>)" |
||||
) |
||||
@ -0,0 +1,286 @@ |
||||
"""Turn ordinary Python functions into OpenAI-compatible tool schemas. |
||||
|
||||
A decorated function's Google-style docstring becomes the tool description the |
||||
model reads, and its type annotations become the JSON schema. That means tool |
||||
documentation lives next to the implementation and cannot drift from it. |
||||
|
||||
@register_tool |
||||
def search(query: str, limit: int = 10) -> str: |
||||
'''Search the speech corpus. |
||||
|
||||
Args: |
||||
query: Words to search for. |
||||
limit: Maximum number of hits. |
||||
''' |
||||
|
||||
Pass ``description=`` to override the docstring — used to load country-specific |
||||
tool prose from ``prompts/tools/*.md`` while keeping the ``Args:`` parsing. |
||||
""" |
||||
from __future__ import annotations |
||||
|
||||
import ast |
||||
import inspect |
||||
import json |
||||
import re |
||||
import types |
||||
from typing import Any, Callable, Dict, Iterable, List, Optional, Union, get_args, get_origin |
||||
|
||||
from pydantic import BaseModel |
||||
|
||||
TOOL_REGISTRY: Dict[str, Dict[str, Any]] = {} |
||||
|
||||
_NoneType = type(None) |
||||
|
||||
|
||||
def _unwrap_optional(annotation: Any) -> tuple[Any, bool]: |
||||
"""Reduce ``Optional[X]`` / ``X | None`` to ``(X, True)``. |
||||
|
||||
Without this, ``get_origin(Optional[list[str]])`` is ``Union`` rather than |
||||
``list``, so a list parameter would be advertised to the model as a string |
||||
and the list coercion in :func:`execute_tool` would never fire. |
||||
""" |
||||
if get_origin(annotation) is Union or isinstance(annotation, types.UnionType): |
||||
args = [a for a in get_args(annotation) if a is not _NoneType] |
||||
optional = len(args) != len(get_args(annotation)) |
||||
if not args: |
||||
return str, optional |
||||
# A genuine multi-type union (e.g. Union[str, List[str]]) has no single |
||||
# JSON type; describe it by its first member, which is what callers coerce to. |
||||
return args[0], optional |
||||
return annotation, False |
||||
|
||||
|
||||
def _pytype_to_jsonschema(annotation: Any) -> dict: |
||||
annotation, _ = _unwrap_optional(annotation) |
||||
|
||||
origin = get_origin(annotation) |
||||
if origin in (list, List): |
||||
args = get_args(annotation) |
||||
return {"type": "array", "items": _pytype_to_jsonschema(args[0] if args else str)} |
||||
|
||||
if inspect.isclass(annotation) and issubclass(annotation, BaseModel): |
||||
return {"type": "object", **annotation.model_json_schema()} |
||||
|
||||
return { |
||||
str: {"type": "string"}, |
||||
int: {"type": "integer"}, |
||||
float: {"type": "number"}, |
||||
bool: {"type": "boolean"}, |
||||
dict: {"type": "object"}, |
||||
list: {"type": "array", "items": {"type": "string"}}, |
||||
}.get(annotation, {"type": "string"}) |
||||
|
||||
|
||||
_SECTION_HEADINGS = frozenset( |
||||
{"returns", "return", "raises", "raise", "yields", "yield", |
||||
"examples", "example", "notes", "note"} |
||||
) |
||||
_PARAM_RE = re.compile(r"^(\w+)\s*(?:\(([^)]+)\))?\s*:\s*(.*)$") |
||||
|
||||
|
||||
def _parse_google_docstring(docstring: Optional[str]) -> dict: |
||||
"""Split a Google-style docstring into a description and per-parameter docs. |
||||
|
||||
Everything outside the ``Args:`` block becomes the description, so ``Returns:`` |
||||
and ``Examples:`` sections still reach the model. |
||||
""" |
||||
if not docstring: |
||||
return {"description": "", "params": {}} |
||||
|
||||
lines = [ln.rstrip() for ln in docstring.splitlines()] |
||||
|
||||
args_start = next( |
||||
(i for i, ln in enumerate(lines) if ln.strip().lower() in ("args:", "arguments:")), |
||||
None, |
||||
) |
||||
args_end = len(lines) |
||||
if args_start is not None: |
||||
for i in range(args_start + 1, len(lines)): |
||||
stripped = lines[i].strip().lower() |
||||
if stripped.endswith(":") and stripped.rstrip(":") in _SECTION_HEADINGS: |
||||
args_end = i |
||||
break |
||||
|
||||
if args_start is None: |
||||
description = " ".join(ln.strip() for ln in lines if ln.strip()) |
||||
return {"description": description.strip(), "params": {}} |
||||
|
||||
desc_parts = [lines[i].strip() for i in range(args_start) if lines[i].strip()] |
||||
desc_parts += [lines[i].strip() for i in range(args_end, len(lines)) if lines[i].strip()] |
||||
|
||||
params: Dict[str, dict] = {} |
||||
i = args_start + 1 |
||||
while i < args_end: |
||||
line = lines[i].strip() |
||||
if not line: |
||||
i += 1 |
||||
continue |
||||
m = _PARAM_RE.match(line) |
||||
if not m: |
||||
i += 1 |
||||
continue |
||||
desc = m.group(3) |
||||
j = i + 1 |
||||
while j < args_end: |
||||
nxt = lines[j].strip() |
||||
if not nxt or _PARAM_RE.match(nxt): |
||||
break |
||||
desc += " " + nxt |
||||
j += 1 |
||||
params[m.group(1)] = {"description": desc.strip(), "type": m.group(2)} |
||||
i = j |
||||
|
||||
return {"description": " ".join(desc_parts).strip(), "params": params} |
||||
|
||||
|
||||
def _openai_function_schema(name: str, description: str, parameters: dict) -> dict: |
||||
params = dict(parameters) |
||||
if params.get("type") != "object": |
||||
params = { |
||||
"type": "object", |
||||
"properties": params.get("properties", params), |
||||
"required": params.get("required", []), |
||||
} |
||||
params.setdefault("additionalProperties", False) |
||||
return { |
||||
"type": "function", |
||||
"function": {"name": name, "description": description, "parameters": params}, |
||||
} |
||||
|
||||
|
||||
def register_tool( |
||||
func: Optional[Callable] = None, |
||||
*, |
||||
name: Optional[str] = None, |
||||
description: Optional[str] = None, |
||||
schema: Optional[dict] = None, |
||||
aliases: Iterable[str] = (), |
||||
): |
||||
"""Register a function as an LLM-callable tool. |
||||
|
||||
Args: |
||||
name: Tool name advertised to the model. Defaults to the function name. |
||||
description: Overrides the docstring description. ``Args:`` parsing still |
||||
applies, so parameter docs keep coming from the docstring. |
||||
schema: Replaces generated parameter schema wholesale. |
||||
aliases: Extra names resolving to the same callable. Excluded from |
||||
:func:`get_tools`, so a renamed tool keeps working when an old name |
||||
is replayed from a persisted conversation. |
||||
""" |
||||
|
||||
def _register(f: Callable) -> Callable: |
||||
fname = name or f.__name__ |
||||
doc = _parse_google_docstring(f.__doc__) |
||||
|
||||
if schema is not None: |
||||
func_schema = schema |
||||
else: |
||||
props, required = {}, [] |
||||
for pname, param in inspect.signature(f).parameters.items(): |
||||
ann = param.annotation if param.annotation is not inspect.Parameter.empty else str |
||||
prop = _pytype_to_jsonschema(ann) |
||||
if pname in doc["params"]: |
||||
prop["description"] = doc["params"][pname]["description"] |
||||
props[pname] = prop |
||||
if param.default is inspect.Parameter.empty: |
||||
required.append(pname) |
||||
func_schema = { |
||||
"type": "object", |
||||
"properties": props, |
||||
"required": required, |
||||
"additionalProperties": False, |
||||
} |
||||
|
||||
entry = { |
||||
"callable": f, |
||||
"schema": _openai_function_schema(fname, description or doc["description"] or "", func_schema), |
||||
"hidden": False, |
||||
} |
||||
TOOL_REGISTRY[fname] = entry |
||||
for alias in aliases: |
||||
TOOL_REGISTRY[alias] = {**entry, "hidden": True} |
||||
return f |
||||
|
||||
return _register if func is None else _register(func) |
||||
|
||||
|
||||
def get_tools( |
||||
specific_tools: Optional[List[str]] = None, |
||||
exclude_tools: Optional[List[str]] = None, |
||||
) -> List[dict]: |
||||
"""Return the OpenAI-format tool list to advertise to a model. |
||||
|
||||
Hidden aliases are never advertised — they exist only so replayed tool calls |
||||
that use a retired name still resolve. |
||||
""" |
||||
if specific_tools and exclude_tools: |
||||
raise ValueError("Pass specific_tools or exclude_tools, not both") |
||||
|
||||
if isinstance(specific_tools, str): |
||||
specific_tools = [specific_tools] |
||||
|
||||
if specific_tools: |
||||
return [TOOL_REGISTRY[t]["schema"] for t in specific_tools if t in TOOL_REGISTRY] |
||||
|
||||
visible = [e["schema"] for e in TOOL_REGISTRY.values() if not e.get("hidden")] |
||||
if exclude_tools: |
||||
excluded = set(exclude_tools) |
||||
return [t for t in visible if t["function"]["name"] not in excluded] |
||||
return visible |
||||
|
||||
|
||||
def parse_function_call_arguments(raw: Any) -> dict: |
||||
"""Best-effort recovery of a tool-call argument payload. |
||||
|
||||
Models sometimes emit not-quite-JSON. Try strict JSON, then Python literals, |
||||
then the first embedded object, before giving up and handing back the raw text |
||||
so the caller can surface a useful error. |
||||
""" |
||||
if isinstance(raw, dict): |
||||
return raw |
||||
if not isinstance(raw, str): |
||||
return {"_raw_unexpected": str(type(raw)), "value": raw} |
||||
|
||||
for parse in (json.loads, ast.literal_eval): |
||||
try: |
||||
parsed = parse(raw) |
||||
if isinstance(parsed, dict): |
||||
return parsed |
||||
except Exception: |
||||
pass |
||||
|
||||
m = re.search(r"\{.*\}", raw, flags=re.DOTALL) |
||||
if m: |
||||
for parse in (json.loads, ast.literal_eval): |
||||
try: |
||||
parsed = parse(m.group(0)) |
||||
if isinstance(parsed, dict): |
||||
return parsed |
||||
except Exception: |
||||
pass |
||||
|
||||
return {"_raw": raw} |
||||
|
||||
|
||||
def execute_tool(name: str, args: dict) -> Any: |
||||
"""Invoke a registered tool, coercing arguments to the declared types.""" |
||||
entry = TOOL_REGISTRY.get(name) |
||||
if not entry: |
||||
raise RuntimeError(f"Tool {name!r} is not registered") |
||||
|
||||
fn = entry["callable"] |
||||
kwargs = {} |
||||
for pname, param in inspect.signature(fn).parameters.items(): |
||||
if pname not in args: |
||||
continue |
||||
val = args[pname] |
||||
ann, _ = _unwrap_optional( |
||||
param.annotation if param.annotation is not inspect.Parameter.empty else None |
||||
) |
||||
# Models frequently send a comma-separated string where a list is declared. |
||||
if get_origin(ann) in (list, List) or ann is list: |
||||
if isinstance(val, str): |
||||
val = [x.strip() for x in val.split(",") if x.strip()] |
||||
kwargs[pname] = val |
||||
|
||||
return fn(**kwargs) |
||||
Loading…
Reference in new issue