feat: add AI catalog description generation
This commit is contained in:
@@ -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"]
|
||||
@@ -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()
|
||||
Generated
+2365
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user