File size: 9,436 Bytes
d61821a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 | """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 {},
)
|