feat: add AI catalog description generation
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Internal process adapters that are not part of the ``tht`` CLI surface."""
|
||||
@@ -0,0 +1,316 @@
|
||||
"""One-shot stdin/stdout adapter around LiteLLM chat completion."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
MAX_STDIN_BYTES = 256 * 1024
|
||||
MAX_API_KEY_BYTES = 16 * 1024
|
||||
MAX_MESSAGES = 32
|
||||
MAX_MESSAGE_CONTENT_BYTES = 64 * 1024
|
||||
MAX_TOTAL_CONTENT_BYTES = 128 * 1024
|
||||
MAX_RESPONSE_CONTENT_BYTES = 64 * 1024
|
||||
MAX_STDERR_BYTES = 256
|
||||
KEYLESS_API_KEY_PLACEHOLDER = "not-required"
|
||||
|
||||
_REQUIRED_KEYS = frozenset({"model", "messages"})
|
||||
_OPTIONAL_KEYS = frozenset({"api_key", "api_base", "api_version", "disable_thinking"})
|
||||
_MESSAGE_KEYS = frozenset({"role", "content"})
|
||||
_MODEL_PATTERN = re.compile(r"[A-Za-z0-9][A-Za-z0-9._:/-]{0,255}\Z")
|
||||
_API_VERSION_PATTERN = re.compile(r"[A-Za-z0-9][A-Za-z0-9._-]{0,127}\Z")
|
||||
_DIAGNOSTICS = {
|
||||
"invalid_request": "completion helper: invalid request\n",
|
||||
"provider_failure": "completion helper: provider failure\n",
|
||||
"invalid_response": "completion helper: invalid response\n",
|
||||
}
|
||||
|
||||
Completion = Callable[..., object]
|
||||
Result = dict[str, object]
|
||||
|
||||
|
||||
class _InvalidRequest(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class _InvalidResponse(Exception):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True, repr=False)
|
||||
class _ValidatedRequest:
|
||||
model: str
|
||||
api_key: str | None
|
||||
messages: tuple[tuple[str, str], ...]
|
||||
api_base: str | None
|
||||
api_version: str | None
|
||||
disable_thinking: bool
|
||||
|
||||
|
||||
def _utf8_size(value: str, error: type[Exception]) -> int:
|
||||
try:
|
||||
return len(value.encode("utf-8"))
|
||||
except UnicodeEncodeError as exc:
|
||||
raise error from exc
|
||||
|
||||
|
||||
def _object_without_duplicates(pairs: list[tuple[str, Any]]) -> dict[str, Any]:
|
||||
value: dict[str, Any] = {}
|
||||
for key, item in pairs:
|
||||
if key in value:
|
||||
raise _InvalidRequest
|
||||
value[key] = item
|
||||
return value
|
||||
|
||||
|
||||
def _reject_nonstandard_number(_: str) -> None:
|
||||
raise _InvalidRequest
|
||||
|
||||
|
||||
def _read_json_request() -> object:
|
||||
raw = sys.stdin.buffer.read(MAX_STDIN_BYTES + 1)
|
||||
if not raw or len(raw) > MAX_STDIN_BYTES:
|
||||
raise _InvalidRequest
|
||||
try:
|
||||
return json.loads(
|
||||
raw.decode("utf-8"),
|
||||
object_pairs_hook=_object_without_duplicates,
|
||||
parse_constant=_reject_nonstandard_number,
|
||||
)
|
||||
except (UnicodeError, ValueError, RecursionError, _InvalidRequest) as exc:
|
||||
raise _InvalidRequest from exc
|
||||
|
||||
|
||||
def _validate_api_base(value: object) -> str:
|
||||
if type(value) is not str or not value or _utf8_size(value, _InvalidRequest) > 2048:
|
||||
raise _InvalidRequest
|
||||
if any(character.isspace() or ord(character) < 32 for character in value):
|
||||
raise _InvalidRequest
|
||||
try:
|
||||
parsed = urlsplit(value)
|
||||
port = parsed.port
|
||||
except ValueError as exc:
|
||||
raise _InvalidRequest from exc
|
||||
if (
|
||||
parsed.scheme.lower() not in {"http", "https"}
|
||||
or not parsed.hostname
|
||||
or parsed.username is not None
|
||||
or parsed.password is not None
|
||||
or parsed.query
|
||||
or parsed.fragment
|
||||
or (port is not None and not 1 <= port <= 65535)
|
||||
):
|
||||
raise _InvalidRequest
|
||||
return value
|
||||
|
||||
|
||||
def _validate_request(value: object) -> _ValidatedRequest:
|
||||
if type(value) is not dict:
|
||||
raise _InvalidRequest
|
||||
keys = set(value)
|
||||
if not _REQUIRED_KEYS.issubset(keys) or not keys.issubset(_REQUIRED_KEYS | _OPTIONAL_KEYS):
|
||||
raise _InvalidRequest
|
||||
|
||||
model = value["model"]
|
||||
if type(model) is not str or _MODEL_PATTERN.fullmatch(model) is None:
|
||||
raise _InvalidRequest
|
||||
|
||||
api_key = None
|
||||
if "api_key" in value:
|
||||
candidate = value["api_key"]
|
||||
if (
|
||||
type(candidate) is not str
|
||||
or not candidate
|
||||
or _utf8_size(candidate, _InvalidRequest) > MAX_API_KEY_BYTES
|
||||
or any(character.isspace() or ord(character) < 32 for character in candidate)
|
||||
):
|
||||
raise _InvalidRequest
|
||||
api_key = candidate
|
||||
|
||||
raw_messages = value["messages"]
|
||||
if type(raw_messages) is not list or not 1 <= len(raw_messages) <= MAX_MESSAGES:
|
||||
raise _InvalidRequest
|
||||
messages: list[tuple[str, str]] = []
|
||||
total_content_bytes = 0
|
||||
for raw_message in raw_messages:
|
||||
if type(raw_message) is not dict or set(raw_message) != _MESSAGE_KEYS:
|
||||
raise _InvalidRequest
|
||||
role = raw_message["role"]
|
||||
content = raw_message["content"]
|
||||
if type(role) is not str or role not in {"system", "user"}:
|
||||
raise _InvalidRequest
|
||||
if type(content) is not str or not content.strip():
|
||||
raise _InvalidRequest
|
||||
content_bytes = _utf8_size(content, _InvalidRequest)
|
||||
if content_bytes > MAX_MESSAGE_CONTENT_BYTES:
|
||||
raise _InvalidRequest
|
||||
total_content_bytes += content_bytes
|
||||
if total_content_bytes > MAX_TOTAL_CONTENT_BYTES:
|
||||
raise _InvalidRequest
|
||||
messages.append((role, content))
|
||||
|
||||
api_base = None
|
||||
if "api_base" in value:
|
||||
api_base = _validate_api_base(value["api_base"])
|
||||
|
||||
api_version = None
|
||||
if "api_version" in value:
|
||||
candidate = value["api_version"]
|
||||
if type(candidate) is not str or _API_VERSION_PATTERN.fullmatch(candidate) is None:
|
||||
raise _InvalidRequest
|
||||
api_version = candidate
|
||||
|
||||
disable_thinking = False
|
||||
if "disable_thinking" in value:
|
||||
if value["disable_thinking"] is not True:
|
||||
raise _InvalidRequest
|
||||
disable_thinking = True
|
||||
|
||||
return _ValidatedRequest(
|
||||
model=model,
|
||||
api_key=api_key,
|
||||
messages=tuple(messages),
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
disable_thinking=disable_thinking,
|
||||
)
|
||||
|
||||
|
||||
def _litellm_completion(**kwargs: object) -> object:
|
||||
import litellm
|
||||
|
||||
litellm.suppress_debug_info = True
|
||||
litellm.set_verbose = False
|
||||
litellm.turn_off_message_logging = True
|
||||
litellm.log_raw_request_response = False
|
||||
litellm.redact_messages_in_exceptions = True
|
||||
litellm.redact_user_api_key_info = True
|
||||
return litellm.completion(**kwargs)
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _silence_provider_output() -> Iterator[None]:
|
||||
with open(os.devnull, "w", encoding="utf-8") as sink:
|
||||
sys.stdout.flush()
|
||||
sys.stderr.flush()
|
||||
saved_stdout = os.dup(1)
|
||||
saved_stderr = os.dup(2)
|
||||
try:
|
||||
os.dup2(sink.fileno(), 1)
|
||||
os.dup2(sink.fileno(), 2)
|
||||
with contextlib.redirect_stdout(sink), contextlib.redirect_stderr(sink):
|
||||
yield
|
||||
finally:
|
||||
os.dup2(saved_stdout, 1)
|
||||
os.dup2(saved_stderr, 2)
|
||||
os.close(saved_stdout)
|
||||
os.close(saved_stderr)
|
||||
|
||||
|
||||
def _provider_kwargs(request: _ValidatedRequest) -> dict[str, object]:
|
||||
kwargs: dict[str, object] = {
|
||||
"model": request.model,
|
||||
# OpenAI-compatible SDKs require a non-empty client value even when the
|
||||
# explicitly configured endpoint does not authenticate requests.
|
||||
"api_key": request.api_key or KEYLESS_API_KEY_PLACEHOLDER,
|
||||
"messages": [
|
||||
{"role": role, "content": content} for role, content in request.messages
|
||||
],
|
||||
"num_retries": 1,
|
||||
"stream": False,
|
||||
}
|
||||
if request.api_base is not None:
|
||||
kwargs["api_base"] = request.api_base
|
||||
if request.api_version is not None:
|
||||
kwargs["api_version"] = request.api_version
|
||||
if request.disable_thinking:
|
||||
kwargs["extra_body"] = {"chat_template_kwargs": {"enable_thinking": False}}
|
||||
return kwargs
|
||||
|
||||
|
||||
def _field(value: object, name: str) -> object:
|
||||
if isinstance(value, Mapping):
|
||||
if name not in value:
|
||||
raise _InvalidResponse
|
||||
return value[name]
|
||||
try:
|
||||
return getattr(value, name)
|
||||
except (AttributeError, TypeError) as exc:
|
||||
raise _InvalidResponse from exc
|
||||
|
||||
|
||||
def _response_content(response: object) -> str:
|
||||
choices = _field(response, "choices")
|
||||
if (
|
||||
not isinstance(choices, Sequence)
|
||||
or isinstance(choices, (str, bytes, bytearray))
|
||||
or len(choices) != 1
|
||||
):
|
||||
raise _InvalidResponse
|
||||
message = _field(choices[0], "message")
|
||||
content = _field(message, "content")
|
||||
if (
|
||||
type(content) is not str
|
||||
or not content.strip()
|
||||
or _utf8_size(content, _InvalidResponse) > MAX_RESPONSE_CONTENT_BYTES
|
||||
):
|
||||
raise _InvalidResponse
|
||||
return content
|
||||
|
||||
|
||||
def handle_request(request: object, *, completion: Completion | None = None) -> Result:
|
||||
"""Validate one request and normalize one injected or LiteLLM completion."""
|
||||
try:
|
||||
validated = _validate_request(request)
|
||||
except _InvalidRequest:
|
||||
return {"ok": False, "error": "invalid_request"}
|
||||
|
||||
provider = completion or _litellm_completion
|
||||
try:
|
||||
with _silence_provider_output():
|
||||
response = provider(**_provider_kwargs(validated))
|
||||
except Exception: # noqa: BLE001 - provider failures cross a redacted process boundary
|
||||
return {"ok": False, "error": "provider_failure"}
|
||||
|
||||
try:
|
||||
with _silence_provider_output():
|
||||
content = _response_content(response)
|
||||
except Exception: # noqa: BLE001 - arbitrary provider objects are untrusted response data
|
||||
return {"ok": False, "error": "invalid_response"}
|
||||
return {"ok": True, "content": content}
|
||||
|
||||
|
||||
def _write_json(payload: Mapping[str, object]) -> None:
|
||||
sys.stdout.write(json.dumps(payload, ensure_ascii=False, separators=(",", ":")) + "\n")
|
||||
|
||||
|
||||
def _write_diagnostic(result: Mapping[str, object]) -> None:
|
||||
error = result.get("error")
|
||||
if not isinstance(error, str):
|
||||
return
|
||||
diagnostic = _DIAGNOSTICS.get(error, "completion helper: failure\n")
|
||||
encoded = diagnostic.encode("utf-8")[:MAX_STDERR_BYTES]
|
||||
sys.stderr.write(encoded.decode("utf-8", errors="ignore"))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
try:
|
||||
request = _read_json_request()
|
||||
except _InvalidRequest:
|
||||
result: Result = {"ok": False, "error": "invalid_request"}
|
||||
else:
|
||||
result = handle_request(request)
|
||||
_write_json(result)
|
||||
if result.get("ok") is False:
|
||||
_write_diagnostic(result)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user