agent-harness / src /agent_harness /lm_studio_embeddings.py
cuber12's picture
Publish agent harness research code and paper artifacts
d61821a verified
Raw
History Blame Contribute Delete
9.44 kB
"""Strict LM Studio client for the pinned dense-retrieval embedding model."""
from __future__ import annotations
from dataclasses import asdict, dataclass
import json
import math
import os
from typing import Any
from urllib.error import HTTPError, URLError
from urllib.request import Request, urlopen
from .specs import EmbeddingSpec
class EmbeddingStudioError(RuntimeError):
"""Raised when embedding discovery, identity checks, or inference fails."""
def validate_embedding_record(spec: EmbeddingSpec, record: dict[str, Any]) -> None:
quantization = record.get("quantization", {})
quantization_name = quantization.get("name") if isinstance(quantization, dict) else quantization
expected = {
"type": "embedding",
"key": spec.model_key,
"display_name": spec.expected_display_name,
"format": spec.expected_format,
"quantization": spec.expected_quantization,
"size_bytes": spec.expected_size_bytes,
"max_context_length": spec.max_context_length,
}
actual = {
"type": record.get("type"),
"key": record.get("key"),
"display_name": record.get("display_name"),
"format": record.get("format"),
"quantization": quantization_name,
"size_bytes": record.get("size_bytes"),
"max_context_length": record.get("max_context_length"),
}
mismatches = [
f"{field}: expected {expected[field]!r}, observed {actual[field]!r}"
for field in expected
if expected[field] != actual[field]
]
loaded_instances = record.get("loaded_instances", [])
if isinstance(loaded_instances, list) and loaded_instances:
loaded_contexts = {
item.get("config", {}).get("context_length")
for item in loaded_instances
if isinstance(item, dict) and isinstance(item.get("config"), dict)
}
if loaded_contexts != {spec.loaded_context_length}:
mismatches.append(
f"loaded context length: expected only {spec.loaded_context_length}, "
f"observed {sorted(loaded_contexts, key=lambda value: str(value))}"
)
if mismatches:
raise EmbeddingStudioError(
"LM Studio embedding runtime does not match EMB001: " + "; ".join(mismatches)
)
@dataclass(frozen=True, slots=True)
class EmbeddingProbeResult:
model_key: str
vector_count: int
vector_dimension: int
l2_norms: tuple[float, ...]
pairwise_cosine: float
usage: dict[str, Any]
def to_dict(self) -> dict[str, Any]:
return asdict(self)
class LMStudioEmbeddingClient:
def __init__(self, spec: EmbeddingSpec, timeout_seconds: float = 30.0):
self.spec = spec
self.timeout_seconds = timeout_seconds
def _headers(self) -> dict[str, str]:
headers = {"Content-Type": "application/json"}
token = os.environ.get(self.spec.api_token_env, "").strip()
if token:
headers["Authorization"] = f"Bearer {token}"
return headers
def _request(
self,
method: str,
endpoint: str,
payload: dict[str, Any] | None = None,
) -> dict[str, Any]:
data = None if payload is None else json.dumps(payload).encode("utf-8")
request = Request(
self.spec.base_url + endpoint,
data=data,
method=method,
headers=self._headers(),
)
try:
with urlopen(request, timeout=self.timeout_seconds) as response:
body = response.read().decode("utf-8")
except HTTPError as exc:
detail = exc.read().decode("utf-8", errors="replace")
raise EmbeddingStudioError(
f"LM Studio returned HTTP {exc.code} for {endpoint}: {detail}"
) from exc
except URLError as exc:
raise EmbeddingStudioError(
f"Cannot connect to LM Studio embeddings at {self.spec.base_url}: {exc.reason}"
) from exc
try:
decoded = json.loads(body)
except json.JSONDecodeError as exc:
raise EmbeddingStudioError(f"LM Studio returned non-JSON data for {endpoint}") from exc
if not isinstance(decoded, dict):
raise EmbeddingStudioError(f"LM Studio returned an unexpected response for {endpoint}")
return decoded
def resolve(self) -> dict[str, Any]:
response = self._request("GET", self.spec.discovery_endpoint)
models = response.get("models", [])
if not isinstance(models, list):
raise EmbeddingStudioError("LM Studio model discovery response has no models list")
matches = [
item
for item in models
if isinstance(item, dict) and item.get("key") == self.spec.model_key
]
if len(matches) != 1:
visible = [
item.get("key")
for item in models
if isinstance(item, dict) and item.get("type") == "embedding"
]
raise EmbeddingStudioError(
f"Expected exactly one {self.spec.model_key!r} record; visible embeddings: {visible}"
)
record = dict(matches[0])
validate_embedding_record(self.spec, record)
return record
def loaded_model_keys(self) -> tuple[str, ...]:
response = self._request("GET", self.spec.discovery_endpoint)
models = response.get("models", [])
if not isinstance(models, list):
raise EmbeddingStudioError("LM Studio model discovery response has no models list")
return tuple(
str(item.get("key"))
for item in models
if isinstance(item, dict)
and isinstance(item.get("loaded_instances"), list)
and item.get("loaded_instances")
)
def embed(self, inputs: list[str]) -> tuple[tuple[float, ...], ...]:
if not inputs or any(not isinstance(item, str) or not item for item in inputs):
raise ValueError("embedding inputs must be non-empty strings")
response = self._request(
"POST",
self.spec.inference_endpoint,
{"model": self.spec.model_key, "input": inputs},
)
return self._vectors_from_response(response, len(inputs))
def _vectors_from_response(
self,
response: dict[str, Any],
expected_count: int,
) -> tuple[tuple[float, ...], ...]:
if response.get("model") != self.spec.model_key:
raise EmbeddingStudioError(
f"Embedding response model mismatch: {response.get('model')!r}"
)
data = response.get("data", [])
if not isinstance(data, list) or len(data) != expected_count:
raise EmbeddingStudioError("Embedding response vector count does not match input count")
ordered = sorted(data, key=lambda item: item.get("index", -1) if isinstance(item, dict) else -1)
vectors: list[tuple[float, ...]] = []
for expected_index, item in enumerate(ordered):
if not isinstance(item, dict) or item.get("index") != expected_index:
raise EmbeddingStudioError("Embedding response indices are malformed")
raw_vector = item.get("embedding")
if not isinstance(raw_vector, list) or len(raw_vector) != self.spec.vector_dimension:
raise EmbeddingStudioError(
f"Expected {self.spec.vector_dimension}-dimensional embedding at index "
f"{expected_index}"
)
try:
vector = tuple(float(value) for value in raw_vector)
except (TypeError, ValueError) as exc:
raise EmbeddingStudioError("Embedding vector contains non-numeric data") from exc
if not all(math.isfinite(value) for value in vector):
raise EmbeddingStudioError("Embedding vector contains non-finite data")
norm = math.sqrt(sum(value * value for value in vector))
if self.spec.normalized and not math.isclose(norm, 1.0, abs_tol=1e-4):
raise EmbeddingStudioError(
f"Embedding at index {expected_index} is not L2-normalized: {norm}"
)
vectors.append(vector)
return tuple(vectors)
def probe(self) -> EmbeddingProbeResult:
inputs = [
"def binary_search(items, target): return -1",
"class HttpClient: pass",
]
response = self._request(
"POST",
self.spec.inference_endpoint,
{"model": self.spec.model_key, "input": inputs},
)
vectors = self._vectors_from_response(response, len(inputs))
if vectors[0] == vectors[1]:
raise EmbeddingStudioError("Distinct probe inputs produced identical vectors")
norms = tuple(math.sqrt(sum(value * value for value in vector)) for vector in vectors)
pairwise_cosine = sum(left * right for left, right in zip(vectors[0], vectors[1]))
usage = response.get("usage", {})
return EmbeddingProbeResult(
model_key=self.spec.model_key,
vector_count=len(vectors),
vector_dimension=self.spec.vector_dimension,
l2_norms=norms,
pairwise_cosine=pairwise_cosine,
usage=dict(usage) if isinstance(usage, dict) else {},
)