feat: add AI catalog description generation

This commit is contained in:
Codex
2026-08-29 16:42:56 +02:00
parent b0afba81ca
commit 376dd5a09d
76 changed files with 14860 additions and 102 deletions
+1
View File
@@ -17,6 +17,7 @@ dependencies = [
"tqdm>=4.66",
"yake>=0.4",
"portalocker>=2.10",
"litellm>=1.98,<2",
]
[project.scripts]
@@ -0,0 +1,421 @@
from __future__ import annotations
import json
import os
import subprocess
import sys
from pathlib import Path
from types import SimpleNamespace
import pytest
from tht.internal.litellm_completion import handle_request
HARNESS_ROOT = Path(__file__).resolve().parents[1]
HELPER_MODULE = "tht.internal.litellm_completion"
def _run_raw_helper(
tmp_path: Path, raw_request: str, fake_litellm: str
) -> subprocess.CompletedProcess[str]:
provider_path = tmp_path / "provider"
provider_path.mkdir()
(provider_path / "litellm.py").write_text(fake_litellm, encoding="utf-8")
env = os.environ.copy()
env["PYTHONPATH"] = os.pathsep.join((str(provider_path), str(HARNESS_ROOT)))
return subprocess.run(
[sys.executable, "-m", HELPER_MODULE],
input=raw_request,
text=True,
capture_output=True,
cwd=HARNESS_ROOT,
env=env,
timeout=10,
check=False,
)
def _run_helper(tmp_path: Path, request: object, fake_litellm: str) -> subprocess.CompletedProcess[str]:
return _run_raw_helper(tmp_path, json.dumps(request), fake_litellm)
def _valid_request() -> dict[str, object]:
return {
"model": "openai/gpt-4.1-mini",
"api_key": "sk-secret",
"messages": [{"role": "user", "content": "private prompt"}],
}
def test_subprocess_emits_one_pristine_success_line_and_silences_provider_output(tmp_path: Path):
secret = "sk-contract-secret"
prompt = "private patient prompt"
request = {
"model": "openai/gpt-4.1-mini",
"api_key": secret,
"messages": [
{"role": "system", "content": "Return one description."},
{"role": "user", "content": prompt},
],
"api_base": "https://models.example.test/v1",
"api_version": "2026-08-01-preview",
}
fake_litellm = f'''
import sys
suppress_debug_info = False
set_verbose = True
turn_off_message_logging = False
log_raw_request_response = True
print("provider import noise")
print("provider import diagnostic", file=sys.stderr)
def completion(**kwargs):
assert suppress_debug_info is True
assert set_verbose is False
assert turn_off_message_logging is True
assert log_raw_request_response is False
assert kwargs == {{
"model": "openai/gpt-4.1-mini",
"api_key": {secret!r},
"messages": [
{{"role": "system", "content": "Return one description."}},
{{"role": "user", "content": {prompt!r}}},
],
"api_base": "https://models.example.test/v1",
"api_version": "2026-08-01-preview",
"num_retries": 1,
"stream": False,
}}
print(repr(kwargs))
print(repr(kwargs), file=sys.stderr)
return {{"choices": [{{"message": {{"content": "Descrizione italiana"}}}}]}}
'''
result = _run_helper(tmp_path, request, fake_litellm)
assert result.returncode == 0
assert result.stdout == '{"ok":true,"content":"Descrizione italiana"}\n'
assert result.stderr == ""
assert secret not in result.stdout + result.stderr
assert prompt not in result.stdout + result.stderr
def test_subprocess_rejects_an_unknown_request_key_before_loading_provider(tmp_path: Path):
request = {**_valid_request(), "unexpected": True}
result = _run_helper(tmp_path, request, 'raise AssertionError("provider was loaded")')
assert result.returncode == 0
assert result.stdout == '{"ok":false,"error":"invalid_request"}\n'
assert result.stderr == "completion helper: invalid request\n"
@pytest.mark.parametrize(
"raw_request",
[
"{not-json",
"[]",
json.dumps({"model": "openai/gpt-4.1-mini", "api_key": "sk-secret"}),
json.dumps({**_valid_request(), "model": "bad model"}),
json.dumps({**_valid_request(), "model": "a" * 257}),
json.dumps({**_valid_request(), "api_key": " "}),
json.dumps({**_valid_request(), "api_key": "x" * (16 * 1024 + 1)}),
json.dumps(
{
**_valid_request(),
"messages": [{"role": "assistant", "content": "not allowed"}],
}
),
json.dumps(
{
**_valid_request(),
"messages": [{"role": "user", "content": "ok", "name": "extra"}],
}
),
json.dumps({**_valid_request(), "messages": []}),
json.dumps(
{
**_valid_request(),
"messages": [{"role": "user", "content": "ok"}] * 33,
}
),
json.dumps(
{
**_valid_request(),
"messages": [{"role": "user"}],
}
),
json.dumps(
{
**_valid_request(),
"messages": [{"role": "user", "content": 1}],
}
),
json.dumps(
{
**_valid_request(),
"messages": [{"role": "user", "content": "x" * (64 * 1024 + 1)}],
}
),
json.dumps(
{
**_valid_request(),
"messages": [
{"role": "user", "content": "x" * (48 * 1024)},
{"role": "system", "content": "y" * (48 * 1024)},
{"role": "user", "content": "z" * (48 * 1024)},
],
}
),
json.dumps({**_valid_request(), "api_base": "file:///tmp/provider"}),
json.dumps({**_valid_request(), "api_base": "https://user@example.test/v1"}),
json.dumps({**_valid_request(), "api_base": "https://example.test/v1?key=value"}),
json.dumps({**_valid_request(), "api_version": "2026 preview"}),
json.dumps({**_valid_request(), "disable_thinking": False}),
(
'{"model":"openai/first","model":"openai/second","api_key":"sk-secret",'
'"messages":[{"role":"user","content":"private prompt"}]}'
),
(
'{"model":"openai/gpt-4.1-mini","api_key":"sk-secret",'
'"messages":[{"role":"user","content":"private prompt"}],"number":'
+ "1" * 5000
+ "}"
),
json.dumps(
{
**_valid_request(),
"messages": [{"role": "user", "content": "x" * (256 * 1024)}],
}
),
],
ids=[
"malformed-json",
"non-object",
"missing-required-key",
"invalid-model",
"long-model",
"empty-secret",
"long-secret",
"invalid-role",
"message-extra-key",
"empty-messages",
"too-many-messages",
"message-missing-key",
"non-string-content",
"long-message",
"aggregate-content-too-long",
"invalid-url-scheme",
"url-credentials",
"url-query",
"invalid-api-version",
"invalid-disable-thinking",
"duplicate-key",
"oversized-number",
"oversized-stdin",
],
)
def test_subprocess_rejects_malformed_or_oversized_requests_without_loading_provider(
tmp_path: Path, raw_request: str
):
result = _run_raw_helper(
tmp_path,
raw_request,
'raise AssertionError("provider was loaded with private prompt and sk-secret")',
)
assert result.returncode == 0
assert result.stdout == '{"ok":false,"error":"invalid_request"}\n'
assert result.stderr == "completion helper: invalid request\n"
assert "private prompt" not in result.stdout + result.stderr
assert "sk-secret" not in result.stdout + result.stderr
def test_subprocess_normalizes_provider_failure_without_leaking_provider_output(tmp_path: Path):
secret = "sk-provider-canary"
prompt = "provider prompt canary"
request = {
**_valid_request(),
"api_key": secret,
"messages": [{"role": "user", "content": prompt}],
}
fake_litellm = '''
import os
import sys
def completion(**kwargs):
leaked = repr(kwargs).encode("utf-8")
print(repr(kwargs))
print(repr(kwargs), file=sys.stderr)
os.write(1, leaked + b"\\n")
os.write(2, leaked + b"\\n")
os.write(1, b"x" * (1024 * 1024))
os.write(2, b"x" * (1024 * 1024))
raise RuntimeError(repr(kwargs))
'''
result = _run_helper(tmp_path, request, fake_litellm)
assert result.returncode == 0
assert result.stdout == '{"ok":false,"error":"provider_failure"}\n'
assert result.stderr == "completion helper: provider failure\n"
assert len(result.stderr.encode("utf-8")) <= 256
assert secret not in result.stdout + result.stderr
assert prompt not in result.stdout + result.stderr
def test_subprocess_normalizes_an_invalid_response_to_one_safe_line(tmp_path: Path):
secret = "sk-invalid-response"
prompt = "invalid response prompt"
request = {
**_valid_request(),
"api_key": secret,
"messages": [{"role": "user", "content": prompt}],
}
fake_litellm = '''
import os
def completion(**kwargs):
os.write(1, repr(kwargs).encode("utf-8"))
os.write(2, repr(kwargs).encode("utf-8"))
return {"choices": [{"message": {"content": ""}}]}
'''
result = _run_helper(tmp_path, request, fake_litellm)
assert result.returncode == 0
assert result.stdout == '{"ok":false,"error":"invalid_response"}\n'
assert result.stderr == "completion helper: invalid response\n"
assert secret not in result.stdout + result.stderr
assert prompt not in result.stdout + result.stderr
def test_subprocess_escapes_multiline_content_without_adding_output_frames(tmp_path: Path):
content = 'Prima riga\nSeconda "riga" — fine'
fake_litellm = f'''
def completion(**kwargs):
return {{"choices": [{{"message": {{"content": {content!r}}}}}]}}
'''
result = _run_helper(tmp_path, _valid_request(), fake_litellm)
assert result.returncode == 0
assert result.stdout == json.dumps(
{"ok": True, "content": content}, ensure_ascii=False, separators=(",", ":")
) + "\n"
assert result.stdout.count("\n") == 1
assert result.stderr == ""
def test_injected_completion_configures_one_retry_without_fallback_or_secret_export():
secret = "sk-injected-canary"
request = {
**_valid_request(),
"api_key": secret,
"messages": [
{"role": "system", "content": "System instruction"},
{"role": "user", "content": "Private prompt"},
],
}
calls: list[dict[str, object]] = []
def completion(**kwargs: object) -> object:
assert secret not in os.environ.values()
calls.append(kwargs)
return SimpleNamespace(
choices=[SimpleNamespace(message=SimpleNamespace(content="Validated content"))]
)
result = handle_request(request, completion=completion)
assert result == {"ok": True, "content": "Validated content"}
assert calls == [{
"model": "openai/gpt-4.1-mini",
"api_key": secret,
"messages": [
{"role": "system", "content": "System instruction"},
{"role": "user", "content": "Private prompt"},
],
"num_retries": 1,
"stream": False,
}]
def test_injected_completion_uses_non_secret_client_placeholder_for_keyless_endpoint():
calls: list[dict[str, object]] = []
def completion(**kwargs: object) -> object:
calls.append(kwargs)
return {"choices": [{"message": {"content": "Descrizione Qwen"}}]}
request = {
"model": "openai/qwen3.6-35b-a3b",
"messages": [{"role": "user", "content": "Invented metadata"}],
"api_base": "https://models.internal.example/v1",
"disable_thinking": True,
}
assert handle_request(request, completion=completion) == {
"ok": True,
"content": "Descrizione Qwen",
}
assert calls == [{
"model": "openai/qwen3.6-35b-a3b",
"api_key": "not-required",
"messages": [{"role": "user", "content": "Invented metadata"}],
"num_retries": 1,
"stream": False,
"api_base": "https://models.internal.example/v1",
"extra_body": {"chat_template_kwargs": {"enable_thinking": False}},
}]
def test_injected_provider_failure_is_not_retried():
attempts = 0
def completion(**kwargs: object) -> object:
nonlocal attempts
attempts += 1
raise RuntimeError(repr(kwargs))
result = handle_request(_valid_request(), completion=completion)
assert result == {"ok": False, "error": "provider_failure"}
assert attempts == 1
@pytest.mark.parametrize(
"response",
[
None,
{},
{"choices": []},
{
"choices": [
{"message": {"content": "first"}},
{"message": {"content": "second"}},
]
},
{"choices": [{"message": {}}]},
{"choices": [{"message": {"content": None}}]},
{"choices": [{"message": {"content": " "}}]},
{"choices": [{"message": {"content": "x" * (64 * 1024 + 1)}}]},
],
ids=[
"none",
"missing-choices",
"empty-choices",
"ambiguous-choices",
"missing-content",
"non-string-content",
"empty-content",
"oversized-content",
],
)
def test_injected_completion_rejects_invalid_response_shapes(response: object):
result = handle_request(_valid_request(), completion=lambda **_: response)
assert result == {"ok": False, "error": "invalid_response"}
@@ -0,0 +1,9 @@
import tomllib
from pathlib import Path
def test_runtime_dependencies_constrain_litellm_to_the_supported_major():
pyproject = Path(__file__).resolve().parents[1] / "pyproject.toml"
project = tomllib.loads(pyproject.read_text(encoding="utf-8"))["project"]
assert "litellm>=1.98,<2" in project["dependencies"]
+1
View File
@@ -0,0 +1 @@
"""Internal process adapters that are not part of the ``tht`` CLI surface."""
+316
View File
@@ -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()
+2365
View File
File diff suppressed because it is too large Load Diff