351 lines
11 KiB
Python
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()
|