"""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()