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