fix(runtime): harden embedding API behavior

This commit is contained in:
Doster-d 2026-07-15 10:43:55 +03:00
parent bf0fa11747
commit fca4533417
4 changed files with 43 additions and 17 deletions

View file

@ -1,4 +1,7 @@
# apps/embedding-runtime/requirements-test.txt
fastapi==0.115.6
fastapi==0.139.0
starlette==1.3.1
pytest==8.3.4
httpx==0.27.2
ruff==0.15.19
ty==0.0.53

View file

@ -1,8 +1,9 @@
# apps/embedding-runtime/requirements.txt
fastapi==0.115.6
uvicorn[standard]==0.34.0
sentence-transformers==3.3.1
transformers==4.47.1
torch==2.5.1
peft==0.14.0
numpy==2.2.1
fastapi==0.139.0
starlette==1.3.1
uvicorn[standard]==0.51.0
sentence-transformers==5.6.0
transformers==5.13.1
torch==2.13.0
peft==0.19.1
numpy==2.5.1

View file

@ -53,7 +53,7 @@ _active_adapter: Optional[str] = None
def _load_model() -> SentenceTransformer:
global _model
if _model is None:
_model = SentenceTransformer(MODEL_NAME, device=DEVICE, trust_remote_code=True)
_model = SentenceTransformer(MODEL_NAME, device=DEVICE, trust_remote_code=False)
return _model
@ -71,11 +71,13 @@ def _apply_adapter(task: Optional[str]) -> None:
return
model = _load_model()
inner = model[0].auto_model # transformers model inside the ST wrapper
if hasattr(inner, "load_adapter"):
load_adapter = getattr(inner, "load_adapter", None)
if callable(load_adapter):
try:
inner.load_adapter(adapter, adapter_name=task)
if hasattr(inner, "set_adapter"):
inner.set_adapter(task)
load_adapter(adapter, adapter_name=task)
set_adapter = getattr(inner, "set_adapter", None)
if callable(set_adapter):
set_adapter(task)
_active_adapter = adapter
except Exception as exc: # noqa: BLE001 — surface a clean 400 to the caller
raise HTTPException(status_code=400, detail=f"adapter load failed for task={task!r}: {exc}")

View file

@ -69,24 +69,24 @@ class FakeCuda:
def build_fake_torch_module() -> ModuleType:
module = ModuleType("torch")
module.cuda = FakeCuda()
setattr(module, "cuda", FakeCuda())
return module
def build_fake_sentence_transformers_module() -> ModuleType:
module = ModuleType("sentence_transformers")
module.SentenceTransformer = FailingSentenceTransformer
setattr(module, "SentenceTransformer", FailingSentenceTransformer)
return module
def build_fake_numpy_module() -> ModuleType:
module = ModuleType("numpy")
module.ndarray = FakeMatrix
setattr(module, "ndarray", FakeMatrix)
def argsort(scores: FakeScores) -> list[int]:
return sorted(range(len(scores.values)), key=lambda index: scores.values[index])
module.argsort = argsort
setattr(module, "argsort", argsort)
return module
@ -100,6 +100,26 @@ def server_module(monkeypatch: pytest.MonkeyPatch) -> ModuleType:
return importlib.import_module("server")
def test_model_loading_disables_remote_code_TASK_8c19d6a7(
server_module: ModuleType, monkeypatch: pytest.MonkeyPatch
) -> None:
# Given: a recording model constructor at the real loader boundary.
constructor_calls: list[tuple[str, str, bool]] = []
class RecordingSentenceTransformer:
def __init__(self, model_name: str, *, device: str, trust_remote_code: bool) -> None:
constructor_calls.append((model_name, device, trust_remote_code))
monkeypatch.setattr(server_module, "SentenceTransformer", RecordingSentenceTransformer)
setattr(server_module, "_model", None)
# When: the runtime lazily loads its configured model.
server_module._load_model()
# Then: repository-supplied model code is never executed.
assert constructor_calls == [(server_module.MODEL_NAME, server_module.DEVICE, False)]
def test_health_remains_available_without_loading_model_SPEC_KSERVE_002_SC_KSERVE_002(
server_module: ModuleType,
) -> None: