limoXD's picture
Release v0.3 local fail-closed sidecar
1e41561 verified
Raw
History Blame Contribute Delete
11.8 kB
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)