fix(runtime): harden embedding API behavior
This commit is contained in:
parent
bf0fa11747
commit
fca4533417
4 changed files with 43 additions and 17 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
12
server.py
12
server.py
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Reference in a new issue