from __future__ import annotations import json import math import queue import subprocess import threading import uuid from collections.abc import Mapping, Sequence from dataclasses import dataclass from typing import BinaryIO from .sidecar import MAX_PROTOCOL_LINE_BYTES, SIDECAR_SCHEMA_VERSION @dataclass(frozen=True, slots=True) class CandidateSource: id: str revision: str @dataclass(frozen=True, slots=True) class SidecarResult: candidates: tuple[Mapping[str, object], ...] changed: bool reason: str request_id: str @dataclass(frozen=True, slots=True) class SidecarControlResult: ok: bool reason: str request_id: str model_loaded: bool class SidecarClient: def __init__( self, command: Sequence[str], *, request_timeout_seconds: float = 0.5, ) -> None: if not command or any(not part for part in command): raise ValueError("command must contain non-empty arguments") if request_timeout_seconds <= 0: raise ValueError("request_timeout_seconds must be positive") self._command = tuple(command) self._request_timeout_seconds = request_timeout_seconds self._process: subprocess.Popen[bytes] | None = None self._responses: queue.Queue[bytes | None] | None = None self._reader: threading.Thread | None = None self._lock = threading.Lock() def __enter__(self) -> SidecarClient: return self def __exit__(self, exc_type: object, exc_value: object, traceback: object) -> None: del exc_type, exc_value, traceback self.close() def rerank( self, *, reading: str, candidates: Sequence[Mapping[str, object]], candidate_source: CandidateSource, mode: str, left_context: Sequence[str] = (), right_context: Sequence[str] = (), ) -> SidecarResult: original = tuple(candidates) request_id = str(uuid.uuid4()) wire_candidates: list[dict[str, object]] = [] for candidate in original: surface = candidate.get("surface") prior_score = candidate.get("prior_score") if ( not isinstance(surface, str) or not isinstance(prior_score, int | float) or isinstance(prior_score, bool) or not math.isfinite(prior_score) ): return SidecarResult(original, False, "invalid_host_input", request_id) wire_candidates.append({"surface": surface, "prior_score": float(prior_score)}) payload: dict[str, object] = { "schema_version": SIDECAR_SCHEMA_VERSION, "request_id": request_id, "op": "rerank", "candidate_source": { "id": candidate_source.id, "revision": candidate_source.revision, }, "mode": mode, "reading": reading, "left_context": list(left_context), "right_context": list(right_context), "candidates": wire_candidates, } with self._lock: response, failure_reason = self._request(payload) return self._validated_rerank_result( original=original, request_id=request_id, response=response, failure_reason=failure_reason, ) def _validated_rerank_result( self, *, original: tuple[Mapping[str, object], ...], request_id: str, response: dict[str, object] | None, failure_reason: str | None, ) -> SidecarResult: if response is None: return SidecarResult( original, False, failure_reason or "sidecar_unavailable", request_id, ) if response.get("schema_version") != SIDECAR_SCHEMA_VERSION: self._discard_process() return SidecarResult(original, False, "invalid_sidecar_response", request_id) if response.get("request_id") != request_id: self._discard_process() return SidecarResult(original, False, "request_id_mismatch", request_id) if response.get("ok") is not True: self._discard_process() return SidecarResult(original, False, "invalid_sidecar_response", request_id) decision = response.get("decision") raw_ranks = response.get("ranked_original_ranks") if not isinstance(decision, dict) or not isinstance(raw_ranks, list): self._discard_process() return SidecarResult(original, False, "invalid_sidecar_response", request_id) changed = decision.get("changed") reason = decision.get("reason") if not isinstance(changed, bool) or not isinstance(reason, str): self._discard_process() return SidecarResult(original, False, "invalid_sidecar_response", request_id) if ( not all(isinstance(rank, int) and not isinstance(rank, bool) for rank in raw_ranks) or sorted(raw_ranks) != list(range(len(original))) or (not changed and raw_ranks != list(range(len(original)))) ): self._discard_process() return SidecarResult(original, False, "invalid_sidecar_response", request_id) return SidecarResult( tuple(original[rank] for rank in raw_ranks), changed, reason, request_id, ) def health(self, *, timeout_seconds: float | None = None) -> SidecarControlResult: return self._control("health", timeout_seconds=timeout_seconds) def warmup(self, *, timeout_seconds: float = 120.0) -> SidecarControlResult: return self._control("warmup", timeout_seconds=timeout_seconds) def close(self) -> None: with self._lock: self._discard_process() def _control( self, operation: str, *, timeout_seconds: float | None ) -> SidecarControlResult: timeout = self._request_timeout_seconds if timeout_seconds is None else timeout_seconds if timeout <= 0: raise ValueError("timeout_seconds must be positive") request_id = str(uuid.uuid4()) payload: dict[str, object] = { "schema_version": SIDECAR_SCHEMA_VERSION, "request_id": request_id, "op": operation, } with self._lock: response, failure_reason = self._request(payload, timeout_seconds=timeout) return self._validated_control_result( operation=operation, request_id=request_id, response=response, failure_reason=failure_reason, ) def _validated_control_result( self, *, operation: str, request_id: str, response: dict[str, object] | None, failure_reason: str | None, ) -> SidecarControlResult: if response is None: return SidecarControlResult( False, failure_reason or "sidecar_unavailable", request_id, False, ) valid = ( response.get("ok") is True and response.get("schema_version") == SIDECAR_SCHEMA_VERSION and response.get("request_id") == request_id and response.get("operation") == operation and isinstance(response.get("model_loaded"), bool) ) if not valid: self._discard_process() return SidecarControlResult(False, "invalid_sidecar_response", request_id, False) model_loaded = response["model_loaded"] assert isinstance(model_loaded, bool) reason = "ready" if operation == "warmup" else "healthy" return SidecarControlResult(True, reason, request_id, model_loaded) def _request( self, payload: dict[str, object], *, timeout_seconds: float | None = None, ) -> tuple[dict[str, object] | None, str | None]: encoded = ( json.dumps(payload, ensure_ascii=False, separators=(",", ":")) + "\n" ).encode("utf-8") if len(encoded) > MAX_PROTOCOL_LINE_BYTES: return None, "request_too_large" try: self._ensure_process() assert self._process is not None assert self._process.stdin is not None assert self._responses is not None self._process.stdin.write(encoded) self._process.stdin.flush() except OSError: self._discard_process() return None, "sidecar_unavailable" try: timeout = ( self._request_timeout_seconds if timeout_seconds is None else timeout_seconds ) raw_response = self._responses.get(timeout=timeout) except queue.Empty: self._discard_process() return None, "sidecar_timeout" if raw_response is None: self._discard_process() return None, "sidecar_exited" if len(raw_response) > MAX_PROTOCOL_LINE_BYTES or not raw_response.endswith(b"\n"): self._discard_process() return None, "invalid_sidecar_response" try: decoded = json.loads(raw_response.decode("utf-8")) if not isinstance(decoded, dict): raise ValueError("response must be an object") return decoded, None except (UnicodeError, ValueError): self._discard_process() return None, "invalid_sidecar_response" def _ensure_process(self) -> None: if self._process is not None and self._process.poll() is None: return self._discard_process() process = subprocess.Popen( self._command, stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=subprocess.DEVNULL, shell=False, ) assert process.stdout is not None responses: queue.Queue[bytes | None] = queue.Queue() reader = threading.Thread( target=self._read_responses, args=(process.stdout, responses), name="deberta-ime-sidecar-reader", daemon=True, ) self._process = process self._responses = responses self._reader = reader reader.start() @staticmethod def _read_responses(stream: BinaryIO, responses: queue.Queue[bytes | None]) -> None: try: while raw_line := stream.readline(MAX_PROTOCOL_LINE_BYTES + 1): responses.put(raw_line) finally: responses.put(None) def _discard_process(self) -> None: process = self._process reader = self._reader self._process = None self._responses = None self._reader = None if process is None: return if process.stdin is not None: try: process.stdin.close() except OSError: pass if process.poll() is None: try: process.terminate() except OSError: pass try: process.wait(timeout=0.5) except subprocess.TimeoutExpired: try: process.kill() except OSError: pass try: process.wait(timeout=1.0) except subprocess.TimeoutExpired: pass if process.stdout is not None: process.stdout.close() if reader is not None and reader is not threading.current_thread(): reader.join(timeout=1.0)