mirror of
https://github.com/mudler/LocalAI.git
synced 2026-07-30 09:57:57 -04:00
fix(distributed): reject wrong-model requests at the backend (#10952) In distributed mode the controller caches a NodeModel row naming a backend's host:port. A worker can recycle a stopped backend's gRPC port for a different model's backend, and probeHealth verifies liveness rather than identity, so the probe succeeds against whatever now occupies the port and the request is dispatched to the wrong backend. The caller gets a silent wrong-model answer. Nothing in the request could catch this: PredictOptions had no model field, so model identity crossed the wire only in ModelOptions.Model at LoadModel time, and the cached-hit path issues no LoadModel. Every backend's "model not loaded" guard checks a nil handle, which a process holding a different model passes, so the stale row was never dropped either. Add PredictOptions.ModelIdentity and enforce it at the point of use: - The controller populates it in gRPCPredictOpts from ModelConfig.Model, the same expression ModelOptions feeds to model.WithModel and therefore the same value the backend received as ModelOptions.Model. Both are read from one config value in one function, so they are equal by construction and the comparison cannot false-reject. - Backends compare it against what they loaded and return NOT_FOUND with a fixed sentinel. Enforced in pkg/grpc/server.go (27 Go backends), an interceptor in backend/python/common (all 36 Python backends, no per-backend change), and the llama-cpp / ik-llama-cpp / ds4 C++ servers. That is every backend with real exposure: kokoros answers all four RPCs with unimplemented and privacy-filter implements none of them. - The router's reconcile drops the stale replica row on a mismatch, so the next request reloads somewhere correct. Empty means "skip the check" on both sides: a controller that predates the field sends nothing, a backend loaded by such a controller has nothing to compare, and the C++ server synthesizes PredictOptions internally for ASR. That keeps upgrades working in both directions. Scoped to the four PredictOptions RPCs. TTSRequest.model and SoundGenerationRequest.model are deliberately NOT validated: FileStagingClient already rewrites them to worker-local absolute paths, so in distributed mode they already differ from the load-time value and comparing them would reject valid requests. IsModelMismatch requires both the NOT_FOUND code and the sentinel, unlike the neighbouring helpers which accept either. insightface's Embedding returns NOT_FOUND "no face detected" on a PredictOptions RPC, and a code-only check would drop a healthy replica row on every faceless image. Assisted-by: Claude Code:claude-opus-4-8[1m] [Read] [Edit] [Bash] Signed-off-by: Ettore Di Giacinto <mudler@localai.io> Co-authored-by: Ettore Di Giacinto <mudler@localai.io>
223 lines
7.7 KiB
Python
223 lines
7.7 KiB
Python
"""Unit tests for model-identity enforcement (model_identity.py).
|
|
|
|
Run inside any backend venv (needs grpcio, which every Python backend has):
|
|
python -m unittest model_identity_test
|
|
|
|
Mirrors the Go coverage in pkg/grpc/model_identity_test.go and
|
|
pkg/grpc/grpcerrors/errors_test.go. The rules under test are the ones whose
|
|
failure modes are silent: enforcement that is wired up but never installed, and
|
|
enforcement that rejects requests it should serve.
|
|
"""
|
|
|
|
import os
|
|
import unittest
|
|
|
|
import grpc
|
|
|
|
import grpc_auth
|
|
import model_identity
|
|
|
|
|
|
class _Aborted(Exception):
|
|
pass
|
|
|
|
|
|
class _FakeContext:
|
|
"""Minimal ServicerContext: abort() records and raises, like the real one."""
|
|
|
|
def __init__(self):
|
|
self.code = None
|
|
self.details = None
|
|
|
|
def abort(self, code, details):
|
|
self.code = code
|
|
self.details = details
|
|
raise _Aborted(details)
|
|
|
|
|
|
class _FakeCallDetails:
|
|
def __init__(self, method):
|
|
self.method = method
|
|
self.invocation_metadata = ()
|
|
|
|
|
|
class _Request:
|
|
"""Stands in for ModelOptions / PredictOptions.
|
|
|
|
The generated protobuf classes are built at container-build time and are
|
|
not importable here, but the interceptor only ever reads two attributes.
|
|
"""
|
|
|
|
def __init__(self, Model="", ModelIdentity=""):
|
|
self.Model = Model
|
|
self.ModelIdentity = ModelIdentity
|
|
|
|
|
|
class _Result:
|
|
def __init__(self, success=True):
|
|
self.success = success
|
|
|
|
|
|
def _handler(behavior, response_streaming=False):
|
|
if response_streaming:
|
|
return grpc.unary_stream_rpc_method_handler(behavior)
|
|
return grpc.unary_unary_rpc_method_handler(behavior)
|
|
|
|
|
|
class TestInterceptorInstalled(unittest.TestCase):
|
|
"""The wiring, which is where this can silently do nothing.
|
|
|
|
get_auth_interceptors() returns early when LOCALAI_GRPC_AUTH_TOKEN is
|
|
unset, which is the DEFAULT configuration. An identity interceptor added
|
|
after that return is never installed on any of the 36 Python backends, and
|
|
nothing else would notice.
|
|
"""
|
|
|
|
def setUp(self):
|
|
self._saved = os.environ.get("LOCALAI_GRPC_AUTH_TOKEN")
|
|
os.environ.pop("LOCALAI_GRPC_AUTH_TOKEN", None)
|
|
|
|
def tearDown(self):
|
|
if self._saved is None:
|
|
os.environ.pop("LOCALAI_GRPC_AUTH_TOKEN", None)
|
|
else:
|
|
os.environ["LOCALAI_GRPC_AUTH_TOKEN"] = self._saved
|
|
|
|
def test_installed_when_auth_is_disabled(self):
|
|
interceptors = grpc_auth.get_auth_interceptors()
|
|
self.assertTrue(
|
|
any(isinstance(i, model_identity.ModelIdentityInterceptor) for i in interceptors),
|
|
"identity enforcement must be installed even with gRPC auth off "
|
|
"(the default); got {!r}".format(interceptors),
|
|
)
|
|
|
|
def test_installed_when_auth_is_disabled_aio(self):
|
|
interceptors = grpc_auth.get_auth_interceptors(aio=True)
|
|
self.assertTrue(
|
|
any(
|
|
isinstance(i, model_identity.AsyncModelIdentityInterceptor)
|
|
for i in interceptors
|
|
),
|
|
"async identity enforcement must be installed with gRPC auth off",
|
|
)
|
|
|
|
def test_installed_alongside_auth_when_enabled(self):
|
|
os.environ["LOCALAI_GRPC_AUTH_TOKEN"] = "secret"
|
|
interceptors = grpc_auth.get_auth_interceptors()
|
|
self.assertTrue(
|
|
any(isinstance(i, model_identity.ModelIdentityInterceptor) for i in interceptors)
|
|
)
|
|
self.assertTrue(
|
|
any(isinstance(i, grpc_auth.TokenAuthInterceptor) for i in interceptors)
|
|
)
|
|
|
|
|
|
class TestMismatchRule(unittest.TestCase):
|
|
"""The pure policy. Every 'serve' case here is a false-rejection guard."""
|
|
|
|
def setUp(self):
|
|
self.state = model_identity.ModelIdentityState()
|
|
|
|
def test_rejects_a_different_model(self):
|
|
self.state.record("a.gguf")
|
|
message = self.state.mismatch("b.gguf")
|
|
self.assertIsNotNone(message)
|
|
self.assertIn(model_identity.MODEL_MISMATCH_SENTINEL, message)
|
|
self.assertIn("a.gguf", message)
|
|
self.assertIn("b.gguf", message)
|
|
|
|
def test_serves_the_same_model(self):
|
|
self.state.record("a.gguf")
|
|
self.assertIsNone(self.state.mismatch("a.gguf"))
|
|
|
|
def test_serves_when_the_request_has_no_identity(self):
|
|
self.state.record("a.gguf")
|
|
self.assertIsNone(self.state.mismatch(""))
|
|
|
|
def test_serves_when_nothing_was_recorded(self):
|
|
self.assertIsNone(self.state.mismatch("b.gguf"))
|
|
|
|
|
|
class TestInterceptorBehavior(unittest.TestCase):
|
|
def setUp(self):
|
|
self.interceptor = model_identity.ModelIdentityInterceptor()
|
|
self.served = []
|
|
|
|
def _intercept(self, method, handler):
|
|
return self.interceptor.intercept_service(
|
|
lambda _: handler, _FakeCallDetails(method)
|
|
)
|
|
|
|
def _load(self, model, success=True):
|
|
handler = _handler(lambda request, context: _Result(success=success))
|
|
wrapped = self._intercept("/backend.Backend/LoadModel", handler)
|
|
wrapped.unary_unary(_Request(Model=model), _FakeContext())
|
|
|
|
def _call(self, method, identity, response_streaming=False):
|
|
def behavior(request, context):
|
|
self.served.append(method)
|
|
return "served"
|
|
|
|
handler = _handler(behavior, response_streaming=response_streaming)
|
|
wrapped = self._intercept(method, handler)
|
|
context = _FakeContext()
|
|
behavior_fn = wrapped.unary_stream if response_streaming else wrapped.unary_unary
|
|
return behavior_fn(_Request(ModelIdentity=identity), context), context
|
|
|
|
def test_load_records_the_identity(self):
|
|
self._load("a.gguf")
|
|
self.assertEqual(self.interceptor.state.loaded, "a.gguf")
|
|
|
|
def test_failed_load_records_nothing(self):
|
|
self._load("a.gguf", success=False)
|
|
self.assertEqual(self.interceptor.state.loaded, "")
|
|
|
|
def test_rejects_every_guarded_rpc_on_mismatch(self):
|
|
self._load("a.gguf")
|
|
for method in sorted(model_identity._GUARDED_METHODS):
|
|
streaming = method.endswith("PredictStream")
|
|
with self.subTest(method=method):
|
|
with self.assertRaises(_Aborted):
|
|
self._call(method, "b.gguf", response_streaming=streaming)
|
|
self.assertEqual(self.served, [], "no request may reach the model")
|
|
|
|
def test_reject_uses_not_found_and_the_sentinel(self):
|
|
self._load("a.gguf")
|
|
|
|
def behavior(request, context):
|
|
return "served"
|
|
|
|
wrapped = self._intercept(
|
|
"/backend.Backend/Predict", _handler(behavior)
|
|
)
|
|
context = _FakeContext()
|
|
with self.assertRaises(_Aborted):
|
|
wrapped.unary_unary(_Request(ModelIdentity="b.gguf"), context)
|
|
self.assertEqual(context.code, grpc.StatusCode.NOT_FOUND)
|
|
self.assertIn(model_identity.MODEL_MISMATCH_SENTINEL, context.details)
|
|
|
|
def test_serves_matching_identity(self):
|
|
self._load("a.gguf")
|
|
for method in sorted(model_identity._GUARDED_METHODS):
|
|
streaming = method.endswith("PredictStream")
|
|
self._call(method, "a.gguf", response_streaming=streaming)
|
|
self.assertEqual(len(self.served), len(model_identity._GUARDED_METHODS))
|
|
|
|
def test_serves_request_without_identity(self):
|
|
self._load("a.gguf")
|
|
self._call("/backend.Backend/Predict", "")
|
|
self.assertEqual(self.served, ["/backend.Backend/Predict"])
|
|
|
|
def test_serves_when_load_recorded_nothing(self):
|
|
self._call("/backend.Backend/Predict", "b.gguf")
|
|
self.assertEqual(self.served, ["/backend.Backend/Predict"])
|
|
|
|
def test_unguarded_rpcs_pass_through_untouched(self):
|
|
handler = _handler(lambda request, context: "served")
|
|
wrapped = self._intercept("/backend.Backend/TTS", handler)
|
|
self.assertIs(wrapped, handler)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|