Files
ThothII/harness/tht/internal/litellm_completion.py

351 lines
11 KiB
Python

"""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 _response_usage(response: object) -> dict[str, int] | None:
"""Normalize LiteLLM usage, tolerating providers that omit it."""
try:
usage = _field(response, "usage")
prompt = _field(usage, "prompt_tokens")
output = _field(usage, "completion_tokens")
if type(prompt) is not int or prompt < 0 or type(output) is not int or output < 0:
return None
cached = 0
for source, name in ((usage, "cache_read_input_tokens"), (usage, "cached_tokens")):
try:
value = _field(source, name)
except _InvalidResponse:
continue
if type(value) is int and value >= 0:
cached = value
break
try:
details = _field(usage, "prompt_tokens_details")
value = _field(details, "cached_tokens")
if type(value) is int and value >= 0:
cached = value
except _InvalidResponse:
pass
cached = min(cached, prompt)
return {"input": prompt - cached, "cacheRead": cached, "output": output}
except _InvalidResponse:
return None
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"}
result: dict[str, object] = {"ok": True, "content": content}
usage = _response_usage(response)
if usage is not None:
result["usage"] = usage
return result
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()