Files
LocalAI/backend/python/common/model_identity.py
mudler's LocalAI [bot] 465d488c90 fix(distributed): reject wrong-model requests at the backend (#10970)
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>
2026-07-20 13:05:47 +02:00

193 lines
7.3 KiB
Python

"""Model-identity enforcement for LocalAI Python backends.
In distributed mode the controller caches a routing row naming a backend's
host:port. A worker can recycle a stopped backend's gRPC port for a different
model's backend, and the controller's health probe checks liveness rather than
identity, so the request is dispatched to whatever now occupies the port and
the caller gets a silent wrong-model answer (#10952).
PredictOptions.ModelIdentity carries the model the request is for, so the
backend can reject it at the point of use. This module enforces that for every
Python backend at once: all 36 of them build their server through
grpc_auth.get_auth_interceptors(), so wiring it there needs no per-backend
change. There is no shared BackendServicer base class to hook instead, and the
backends store their loaded model in wildly different attributes, so an
interceptor is the only single point that sees both the LoadModel request and
the inference requests.
Enforcement is deliberately narrow: it compares two strings and never inspects
the model itself.
"""
import threading
import grpc
# Must match grpcerrors.ModelMismatchSentinel in pkg/grpc/grpcerrors/errors.go.
# The router requires this substring AND the NOT_FOUND code before it treats a
# reply as a mismatch, because NOT_FOUND alone is not exclusively ours on these
# RPCs (insightface's Embedding returns it for "no face detected").
MODEL_MISMATCH_SENTINEL = "model identity mismatch"
_LOAD_METHOD = "/backend.Backend/LoadModel"
# The four RPCs that carry PredictOptions. Nothing else has an identity field.
# TTS and SoundGeneration are excluded on purpose: their `model` field is
# already rewritten to a worker-local path by the controller's
# FileStagingClient, so comparing it would reject valid requests.
_GUARDED_METHODS = frozenset(
(
"/backend.Backend/Predict",
"/backend.Backend/PredictStream",
"/backend.Backend/Embedding",
"/backend.Backend/TokenizeString",
)
)
class ModelIdentityState:
"""The identity this process loaded, and the rule for judging a request.
A backend process serves exactly one model (worker process keys are
model+backend+replica), so a single value is enough.
"""
def __init__(self):
self._lock = threading.Lock()
self._loaded = ""
def record(self, model: str) -> None:
with self._lock:
self._loaded = model or ""
@property
def loaded(self) -> str:
with self._lock:
return self._loaded
def mismatch(self, requested: str):
"""Return an error message when `requested` names another model.
Either side being empty means "skip": the request side is empty for a
controller that predates the field and for internally synthesized
requests, and the loaded side is empty when such a controller performed
the load. Neither can judge the other, and a false rejection is worse
than the miss it prevents.
"""
if not requested:
return None
loaded = self.loaded
if not loaded or loaded == requested:
return None
return "{}: loaded {!r}, requested {!r}".format(
MODEL_MISMATCH_SENTINEL, loaded, requested
)
def _rebuild(handler, behavior):
"""Return a copy of `handler` with its behavior replaced.
Only unary-request handlers are ever passed here: LoadModel and the four
guarded RPCs all take a single request message.
"""
if handler.response_streaming:
return grpc.unary_stream_rpc_method_handler(
behavior,
request_deserializer=handler.request_deserializer,
response_serializer=handler.response_serializer,
)
return grpc.unary_unary_rpc_method_handler(
behavior,
request_deserializer=handler.request_deserializer,
response_serializer=handler.response_serializer,
)
class ModelIdentityInterceptor(grpc.ServerInterceptor):
"""Sync interceptor that records the loaded model and guards inference."""
def __init__(self, state: ModelIdentityState = None):
self.state = state or ModelIdentityState()
def intercept_service(self, continuation, handler_call_details):
method = handler_call_details.method
if method != _LOAD_METHOD and method not in _GUARDED_METHODS:
return continuation(handler_call_details)
handler = continuation(handler_call_details)
if handler is None:
return handler
if method == _LOAD_METHOD:
original = handler.unary_unary
def record(request, context):
result = original(request, context)
# Only a successful load owns the identity; a failed one leaves
# no model, which the model-not-loaded signal already covers.
if getattr(result, "success", True):
self.state.record(getattr(request, "Model", ""))
return result
return _rebuild(handler, record)
original = handler.unary_stream if handler.response_streaming else handler.unary_unary
def guard(request, context):
message = self.state.mismatch(getattr(request, "ModelIdentity", ""))
if message is not None:
# abort() raises, so the request never reaches the model.
context.abort(grpc.StatusCode.NOT_FOUND, message)
return original(request, context)
return _rebuild(handler, guard)
class AsyncModelIdentityInterceptor(grpc.aio.ServerInterceptor):
"""Async counterpart for backends running grpc.aio servers."""
def __init__(self, state: ModelIdentityState = None):
self.state = state or ModelIdentityState()
async def intercept_service(self, continuation, handler_call_details):
method = handler_call_details.method
if method != _LOAD_METHOD and method not in _GUARDED_METHODS:
return await continuation(handler_call_details)
handler = await continuation(handler_call_details)
if handler is None:
return handler
if method == _LOAD_METHOD:
original = handler.unary_unary
async def record(request, context):
result = await original(request, context)
if getattr(result, "success", True):
self.state.record(getattr(request, "Model", ""))
return result
return _rebuild(handler, record)
if handler.response_streaming:
original_stream = handler.unary_stream
async def guard_stream(request, context):
message = self.state.mismatch(getattr(request, "ModelIdentity", ""))
if message is not None:
await context.abort(grpc.StatusCode.NOT_FOUND, message)
async for response in original_stream(request, context):
yield response
return _rebuild(handler, guard_stream)
original_unary = handler.unary_unary
async def guard(request, context):
message = self.state.mismatch(getattr(request, "ModelIdentity", ""))
if message is not None:
await context.abort(grpc.StatusCode.NOT_FOUND, message)
return await original_unary(request, context)
return _rebuild(handler, guard)