From ba2d6d608952e20999eae80ad89fa9d2609323c3 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Mon, 20 Jul 2026 18:32:27 +0000 Subject: [PATCH] fix(distributed): reject wrong-model requests on the remaining modalities #10970 gave the four PredictOptions RPCs a model-identity check so a backend reached through a stale distributed route rejects the request instead of answering from whatever model it holds (#10952). Every other modality shares that exposure: the route is cached by host:port, a worker can recycle a stopped backend's port for another model's backend, and a liveness-only probe cannot tell a stale row from a valid one. Extends the same mechanism to the 21 remaining request messages that reach a backend through the router, using the pattern #10970 established rather than a parallel one: - proto: ModelIdentity on each modality request message. - controller: populated from ModelConfig.Model at the call site that also builds ModelOptions, so load-time and request-time values are equal by construction. - backends: one generic guard in pkg/grpc/server.go (27 Go backends), the method set in backend/python/common (36 Python backends), llama-cpp (AudioTranscription/Stream, Rerank, Score) and privacy-filter (TokenClassify). - reconcile already drops the stale row on IsModelMismatch; no change. TTSRequest and SoundGenerationRequest get a SEPARATE ModelIdentity field rather than reusing their existing `model`: FileStagingClient rewrites `model` to a worker-local path, so comparing it would reject valid requests in exactly the configuration this guards. AudioEncode/AudioDecode are deliberately left unguarded: the opus codec backend is loaded from a literal rather than a ModelConfig, so no value carries the equality guarantee the comparison depends on. The four bidirectional stream RPCs are out of scope; they bypass reconcile. Empty means skip on both sides, so an old controller, an old backend, and the bare request structs in tests/e2e-backends all keep working. Signed-off-by: Ettore Di Giacinto Assisted-by: Claude Code:claude-opus-4-8 [Read] [Edit] [Bash] --- backend/backend.proto | 101 ++++- backend/cpp/llama-cpp/grpc-server.cpp | 14 +- backend/cpp/privacy-filter/grpc-server.cpp | 36 ++ backend/python/common/model_identity.py | 52 ++- backend/python/common/model_identity_test.py | 88 ++++- core/backend/audio_transform.go | 1 + core/backend/depth.go | 4 + core/backend/detection.go | 11 +- core/backend/diarization.go | 24 +- core/backend/face_analyze.go | 7 +- core/backend/face_verify.go | 9 +- core/backend/image.go | 5 + .../backend/model_identity_modalities_test.go | 346 ++++++++++++++++++ core/backend/rerank.go | 5 + core/backend/score.go | 1 + core/backend/sound_classification.go | 12 +- core/backend/soundgeneration.go | 21 +- core/backend/token_classify.go | 5 +- core/backend/transcript.go | 12 +- core/backend/tts.go | 27 +- core/backend/tts_test.go | 27 +- core/backend/vad.go | 3 +- core/backend/video.go | 3 + core/backend/voice_analyze.go | 5 +- core/backend/voice_embed.go | 3 +- core/backend/voice_verify.go | 9 +- pkg/grpc/model_identity_modalities_test.go | 226 ++++++++++++ pkg/grpc/server.go | 80 +++- tests/e2e/mock-backend/main.go | 52 ++- 29 files changed, 1104 insertions(+), 85 deletions(-) create mode 100644 core/backend/model_identity_modalities_test.go create mode 100644 pkg/grpc/model_identity_modalities_test.go diff --git a/backend/backend.proto b/backend/backend.proto index 107e2a94e1bd..25f45eac89cf 100644 --- a/backend/backend.proto +++ b/backend/backend.proto @@ -136,6 +136,10 @@ message MetricsResponse { message TokenClassifyRequest { string text = 1; float threshold = 2; + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 3; } // TokenClassifyEntity is one detected entity span. Byte offsets are @@ -173,6 +177,10 @@ message ScoreRequest { // candidates differ in length and the consumer wants a per-token // measure comparable across them (PMI-style scoring). bool length_normalize = 4; + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 5; } // CandidateScore is one row in the ScoreResponse, matching by index @@ -204,6 +212,10 @@ message RerankRequest { string query = 1; repeated string documents = 2; int32 top_n = 3; + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 4; } message RerankResult { @@ -336,7 +348,14 @@ message PredictOptions { // paths (core/services/nodes/file_staging_client.go), so in distributed mode // they already differ from the load-time value and comparing them would // reject valid requests. Extending identity to those RPCs needs a separate - // field carrying the untranslated value. + // field carrying the untranslated value - which is exactly what + // TTSRequest.ModelIdentity and SoundGenerationRequest.ModelIdentity are. + // + // Every other request message that reaches a backend through the distributed + // router now carries the same ModelIdentity field, populated from the same + // ModelConfig.Model. FileStagingClient rewrites Src/Dst/Voice/Model/ + // StartImage/EndImage/Audio and never ModelIdentity, so what the backend + // compares is always what the controller sent. string ModelIdentity = 54; // 24 was never assigned; reserve it so it is not silently reused. @@ -510,6 +529,10 @@ message TranscriptRequest { float temperature = 8; repeated string timestamp_granularities = 9; bool stream = 10; + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 11; } message TranscriptResult { @@ -588,6 +611,10 @@ message GenerateImageRequest { // Reference images for models that support them (e.g., Flux Kontext) repeated string ref_images = 12; + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 13; } message GenerateVideoRequest { @@ -607,6 +634,10 @@ message GenerateVideoRequest { // Backend-specific per-request generation parameters. Values are strings // and are validated/coerced by the selected backend. map params = 14; + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 15; } message TTSRequest { @@ -624,10 +655,26 @@ message TTSRequest { // (e.g. Chatterbox exaggeration/cfg_weight/temperature). Values are strings and // coerced by the backend; unset leaves the backend's configured defaults. map params = 7; + // ModelIdentity is a SEPARATE field from `model` above and carries the + // UNTRANSLATED controller-side ModelConfig.Model, so a backend can reject a + // request that reached it through a stale distributed route (#10952). + // + // `model` cannot be reused for this: FileStagingClient.TTS/.TTSStream and the + // SoundGeneration path rewrite it into a worker-local absolute path + // (core/services/nodes/file_staging_client.go), while the load-time value is + // untranslated. In distributed mode - exactly the configuration this guards - + // the two already differ, so comparing them would reject valid requests. + // + // Empty means "no identity supplied" and backends MUST skip the check. + string ModelIdentity = 8; } message VADRequest { repeated float audio = 1; + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 2; } message VADSegment { @@ -659,6 +706,10 @@ message DiarizeRequest { float min_duration_on = 8; // discard segments shorter than this (seconds); 0 = backend default float min_duration_off = 9; // merge gaps shorter than this (seconds); 0 = backend default bool include_text = 10; // when the backend can emit per-segment transcript for free, ask it to populate `text` + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 11; } message DiarizeSegment { @@ -693,6 +744,18 @@ message SoundGenerationRequest { optional string language = 14; optional string timesignature = 15; optional bool instrumental = 17; + // ModelIdentity is a SEPARATE field from `model` above and carries the + // UNTRANSLATED controller-side ModelConfig.Model, so a backend can reject a + // request that reached it through a stale distributed route (#10952). + // + // `model` cannot be reused for this: FileStagingClient.TTS/.TTSStream and the + // SoundGeneration path rewrite it into a worker-local absolute path + // (core/services/nodes/file_staging_client.go), while the load-time value is + // untranslated. In distributed mode - exactly the configuration this guards - + // the two already differ, so comparing them would reject valid requests. + // + // Empty means "no identity supplied" and backends MUST skip the check. + string ModelIdentity = 18; } message TokenizationResponse { @@ -732,6 +795,10 @@ message DetectOptions { repeated float points = 3; // Point coordinates as [x1, y1, label1, x2, y2, label2, ...] (label: 1=pos, 0=neg) repeated float boxes = 4; // Box coordinates as [x1, y1, x2, y2, ...] float threshold = 5; // Detection confidence threshold + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 6; } message Detection { @@ -754,6 +821,10 @@ message SoundDetectionRequest { string src = 1; // audio file path (LocalAI writes the upload to disk) int32 top_k = 2; // number of top tags to return (0 = all classes) float threshold = 3; // optional: drop tags scoring below this + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 4; } message SoundClass { @@ -778,6 +849,10 @@ message DepthRequest { bool include_points = 7; // back-project to a 3D point cloud (DualDPT) float points_conf_thresh = 8; // keep points with confidence >= this threshold repeated string exports = 9; // requested exports: "glb", "colmap" + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 10; } message DepthResponse { @@ -809,6 +884,10 @@ message FaceVerifyRequest { string img2 = 2; // base64-encoded image float threshold = 3; // cosine-distance threshold; 0 = use backend default bool anti_spoofing = 4; // run MiniFASNet liveness on each image; failed liveness forces verified=false + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 5; } message FaceVerifyResponse { @@ -830,6 +909,10 @@ message FaceAnalyzeRequest { string img = 1; // base64-encoded image repeated string actions = 2; // subset of ["age","gender","emotion","race"]; empty = all-supported bool anti_spoofing = 3; + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 4; } message FaceAnalysis { @@ -862,6 +945,10 @@ message VoiceVerifyRequest { string audio2 = 2; // path to second audio clip float threshold = 3; // cosine-distance threshold; 0 = use backend default bool anti_spoofing = 4; // reserved for future AASIST bolt-on + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 5; } message VoiceVerifyResponse { @@ -876,6 +963,10 @@ message VoiceVerifyResponse { message VoiceAnalyzeRequest { string audio = 1; // path to audio clip repeated string actions = 2; // subset of ["age","gender","emotion"]; empty = all-supported + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 3; } message VoiceAnalysis { @@ -894,6 +985,10 @@ message VoiceAnalyzeResponse { message VoiceEmbedRequest { string audio = 1; // path to audio clip + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 2; } message VoiceEmbedResponse { @@ -988,6 +1083,10 @@ message AudioTransformRequest { string reference_path = 2; // optional auxiliary; empty => zero-fill string dst = 3; // required, output file path map params = 4; // backend-specific tuning + // ModelIdentity names the model this request is for; see + // PredictOptions.ModelIdentity for the full rationale. Empty means "no + // identity supplied" and backends MUST skip the check. + string ModelIdentity = 5; } message AudioTransformResult { diff --git a/backend/cpp/llama-cpp/grpc-server.cpp b/backend/cpp/llama-cpp/grpc-server.cpp index 995b61affd8c..88b47a654cbf 100644 --- a/backend/cpp/llama-cpp/grpc-server.cpp +++ b/backend/cpp/llama-cpp/grpc-server.cpp @@ -1416,7 +1416,11 @@ class BackendServiceImpl final : public backend::Backend::Service { // field and for the synthetic PredictOptions this server builds internally // for ASR, and the loaded side is empty when such a controller performed // the load. A false rejection is worse than the miss it prevents. - grpc::Status checkModelIdentity(const backend::PredictOptions* request) { + // Templated over the request type: every guarded request message exposes + // modelidentity(), and one body keeps the rule identical across modalities + // rather than repeating it per RPC. + template + grpc::Status checkModelIdentity(const Request* request) { if (request == nullptr || request->modelidentity().empty()) { return grpc::Status::OK; } @@ -2846,6 +2850,8 @@ class BackendServiceImpl final : public backend::Backend::Service { } grpc::Status Rerank(ServerContext* context, const backend::RerankRequest* request, backend::RerankResult* rerankResult) override { + auto identity = checkModelIdentity(request); + if (!identity.ok()) return identity; if (!params_base.embedding || params_base.pooling_type != LLAMA_POOLING_TYPE_RANK) { return grpc::Status(grpc::StatusCode::UNIMPLEMENTED, "This server does not support reranking. Start it with `--reranking` and without `--embedding`"); } @@ -2970,6 +2976,8 @@ class BackendServiceImpl final : public backend::Backend::Service { grpc::Status Score(ServerContext* context, const backend::ScoreRequest* request, backend::ScoreResponse* response) override { auto auth = checkAuth(context); if (!auth.ok()) return auth; + auto identity = checkModelIdentity(request); + if (!identity.ok()) return identity; if (params_base.model.path.empty()) { return grpc::Status(grpc::StatusCode::FAILED_PRECONDITION, "Model not loaded"); } @@ -3428,6 +3436,8 @@ class BackendServiceImpl final : public backend::Backend::Service { backend::TranscriptResult* response) override { auto auth = checkAuth(context); if (!auth.ok()) return auth; + auto identity = checkModelIdentity(request); + if (!identity.ok()) return identity; backend::Reply reply; grpc::Status st = runTranscriptionAsCompletion(context, request, &reply); @@ -3446,6 +3456,8 @@ class BackendServiceImpl final : public backend::Backend::Service { grpc::ServerWriter* writer) override { auto auth = checkAuth(context); if (!auth.ok()) return auth; + auto identity = checkModelIdentity(request); + if (!identity.ok()) return identity; // Buffered streaming: run the transcription as a normal chat // completion, then emit one delta + one final event. Real diff --git a/backend/cpp/privacy-filter/grpc-server.cpp b/backend/cpp/privacy-filter/grpc-server.cpp index 5b3f7865801f..0221d1b95492 100644 --- a/backend/cpp/privacy-filter/grpc-server.cpp +++ b/backend/cpp/privacy-filter/grpc-server.cpp @@ -41,6 +41,11 @@ namespace { // per loaded model. g_mu guards (re)load against in-flight classification. std::mutex g_mu; pf_ctx * g_ctx = nullptr; +// The ModelOptions.Model this process loaded, compared against +// TokenClassifyRequest.ModelIdentity so a request that arrived through a stale +// distributed route is rejected rather than answered from the wrong model +// (#10952). Guarded by g_mu like the rest of the engine state. +std::string g_loaded_model_identity; std::atomic g_server{nullptr}; // Resolve the device string the engine expects ("cpu" / "gpu" / "cuda" / @@ -113,17 +118,48 @@ class PrivacyFilterBackend final : public backend::Backend::Service { } g_ctx = ctx; + // Record what we loaded so TokenClassify can reject a request meant + // for a different model. request->model(), not modelfile(): it is the + // value the controller also sends as ModelIdentity, and the two are + // read from the same ModelConfig.Model (#10952). + g_loaded_model_identity = request->model(); result->set_success(true); result->set_message("privacy-filter loaded (" + device + ")"); return GStatus::OK; } + // checkModelIdentity mirrors pkg/grpc/server.go, + // backend/python/common/model_identity.py and the llama-cpp server. In + // distributed mode a worker can recycle a stopped backend's gRPC port for + // another model's backend, and the controller's liveness-only probe cannot + // tell a stale cached route from a valid one, so the backend has to catch + // it. Either side empty means "skip": the request side is empty for a + // controller that predates the field, the loaded side when such a + // controller performed the load. A false rejection is worse than the miss. + // Callers must already hold g_mu. + GStatus checkModelIdentity(const backend::TokenClassifyRequest * request) { + if (request == nullptr || request->modelidentity().empty()) { + return GStatus::OK; + } + if (g_loaded_model_identity.empty() || + g_loaded_model_identity == request->modelidentity()) { + return GStatus::OK; + } + // NOT_FOUND plus this exact sentinel is the cross-language contract + // the router matches on (grpcerrors.ModelMismatchSentinel). + return GStatus(StatusCode::NOT_FOUND, + "privacy-filter: model identity mismatch: loaded \"" + + g_loaded_model_identity + "\", requested \"" + + request->modelidentity() + "\""); + } + GStatus TokenClassify(ServerContext *, const backend::TokenClassifyRequest * request, backend::TokenClassifyResponse * response) override { std::lock_guard lock(g_mu); if (!g_ctx) { return GStatus(StatusCode::FAILED_PRECONDITION, "Model not loaded"); } + if (GStatus id = checkModelIdentity(request); !id.ok()) return id; const std::string & text = request->text(); if (text.empty()) { diff --git a/backend/python/common/model_identity.py b/backend/python/common/model_identity.py index fb64f93add76..256b73bb871e 100644 --- a/backend/python/common/model_identity.py +++ b/backend/python/common/model_identity.py @@ -6,7 +6,8 @@ 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 +Every request message that reaches a backend through the distributed router +carries a ModelIdentity field naming 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 @@ -31,16 +32,51 @@ _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. +# Every RPC whose request message carries a ModelIdentity field. This set IS +# the enforcement surface for all 36 Python backends: an RPC missing here is +# silently unprotected, so model_identity_test.py pins the full list. +# +# The guard reads request.ModelIdentity generically, so nothing here is +# modality-specific — a backend that does not implement an RPC simply never +# sees it. +# +# TTS and SoundGeneration are guarded on ModelIdentity, NOT on their `model` +# field: the controller's FileStagingClient rewrites `model` to a worker-local +# absolute path, so comparing that would reject valid requests in distributed +# mode. ModelIdentity is a separate, untranslated field for exactly that reason. +# +# AudioEncode/AudioDecode are absent deliberately: the opus codec backend they +# target is loaded from a literal rather than a ModelConfig, so no value carries +# the load-time/request-time equality guarantee this comparison depends on. _GUARDED_METHODS = frozenset( ( + # PredictOptions RPCs (#10970) "/backend.Backend/Predict", "/backend.Backend/PredictStream", "/backend.Backend/Embedding", "/backend.Backend/TokenizeString", + # Remaining modalities + "/backend.Backend/GenerateImage", + "/backend.Backend/GenerateVideo", + "/backend.Backend/TTS", + "/backend.Backend/TTSStream", + "/backend.Backend/SoundGeneration", + "/backend.Backend/AudioTranscription", + "/backend.Backend/AudioTranscriptionStream", + "/backend.Backend/Detect", + "/backend.Backend/Depth", + "/backend.Backend/FaceVerify", + "/backend.Backend/FaceAnalyze", + "/backend.Backend/VoiceVerify", + "/backend.Backend/VoiceAnalyze", + "/backend.Backend/VoiceEmbed", + "/backend.Backend/Rerank", + "/backend.Backend/TokenClassify", + "/backend.Backend/Score", + "/backend.Backend/VAD", + "/backend.Backend/Diarize", + "/backend.Backend/SoundDetection", + "/backend.Backend/AudioTransform", ) ) @@ -87,8 +123,10 @@ def mismatch(self, requested: str): 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. + Only unary-request handlers are ever passed here: LoadModel and every + entry in _GUARDED_METHODS take a single request message. The bidirectional + streams (AudioTranscriptionLive, AudioTransformStream, AudioToAudioStream, + Forward) are not guarded and never reach this function. """ if handler.response_streaming: return grpc.unary_stream_rpc_method_handler( diff --git a/backend/python/common/model_identity_test.py b/backend/python/common/model_identity_test.py index 2cb7011170b8..7d3019007eb5 100644 --- a/backend/python/common/model_identity_test.py +++ b/backend/python/common/model_identity_test.py @@ -214,9 +214,95 @@ def test_serves_when_load_recorded_nothing(self): def test_unguarded_rpcs_pass_through_untouched(self): handler = _handler(lambda request, context: "served") - wrapped = self._intercept("/backend.Backend/TTS", handler) + wrapped = self._intercept("/backend.Backend/Health", handler) self.assertIs(wrapped, handler) +# Every modality request message now carries a ModelIdentity field, so every +# modality RPC shares the guard. The set below is the enforcement surface for +# all 36 Python backends at once: an RPC missing from it is silently +# unprotected, which is the failure mode this class exists to catch. +_EXPECTED_MODALITY_METHODS = ( + "/backend.Backend/GenerateImage", + "/backend.Backend/GenerateVideo", + "/backend.Backend/TTS", + "/backend.Backend/TTSStream", + "/backend.Backend/SoundGeneration", + "/backend.Backend/AudioTranscription", + "/backend.Backend/AudioTranscriptionStream", + "/backend.Backend/Detect", + "/backend.Backend/Depth", + "/backend.Backend/FaceVerify", + "/backend.Backend/FaceAnalyze", + "/backend.Backend/VoiceVerify", + "/backend.Backend/VoiceAnalyze", + "/backend.Backend/VoiceEmbed", + "/backend.Backend/Rerank", + "/backend.Backend/TokenClassify", + "/backend.Backend/Score", + "/backend.Backend/VAD", + "/backend.Backend/Diarize", + "/backend.Backend/SoundDetection", + "/backend.Backend/AudioTransform", +) + + +class TestModalityMethods(unittest.TestCase): + def setUp(self): + self.interceptor = model_identity.ModelIdentityInterceptor() + self.interceptor.state.record("a.gguf") + self.served = [] + + 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.interceptor.intercept_service( + lambda _: handler, _FakeCallDetails(method) + ) + fn = wrapped.unary_stream if response_streaming else wrapped.unary_unary + return fn(_Request(ModelIdentity=identity), _FakeContext()) + + def test_every_modality_rpc_is_guarded(self): + for method in _EXPECTED_MODALITY_METHODS: + with self.subTest(method=method): + self.assertIn( + method, + model_identity._GUARDED_METHODS, + "{} is unprotected on all Python backends".format(method), + ) + + def test_every_modality_rpc_rejects_a_mismatch(self): + for method in _EXPECTED_MODALITY_METHODS: + streaming = method.endswith("Stream") + 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_every_modality_rpc_serves_a_match(self): + for method in _EXPECTED_MODALITY_METHODS: + streaming = method.endswith("Stream") + self._call(method, "a.gguf", response_streaming=streaming) + self.assertEqual(len(self.served), len(_EXPECTED_MODALITY_METHODS)) + + # Compatibility: an old controller sends nothing, and the e2e backend suite + # drives real backends with bare request structs. + def test_every_modality_rpc_serves_without_an_identity(self): + for method in _EXPECTED_MODALITY_METHODS: + streaming = method.endswith("Stream") + self._call(method, "", response_streaming=streaming) + self.assertEqual(len(self.served), len(_EXPECTED_MODALITY_METHODS)) + + # AudioEncode/AudioDecode stay out: the opus codec backend is loaded from a + # literal, not a ModelConfig, so no value carries the structural guarantee + # the comparison depends on. + def test_codec_rpcs_stay_unguarded(self): + for method in ("/backend.Backend/AudioEncode", "/backend.Backend/AudioDecode"): + self.assertNotIn(method, model_identity._GUARDED_METHODS) + + if __name__ == "__main__": unittest.main() diff --git a/core/backend/audio_transform.go b/core/backend/audio_transform.go index 399e5c803eac..bb006fc1077c 100644 --- a/core/backend/audio_transform.go +++ b/core/backend/audio_transform.go @@ -83,6 +83,7 @@ func ModelAudioTransform( } res, err := transformModel.AudioTransform(ctx, &proto.AudioTransformRequest{ + ModelIdentity: modelConfig.Model, AudioPath: audioPath, ReferencePath: referencePath, Dst: dst, diff --git a/core/backend/depth.go b/core/backend/depth.go index ca41ae233cf5..154cf3335ef1 100644 --- a/core/backend/depth.go +++ b/core/backend/depth.go @@ -40,6 +40,10 @@ func Depth( startTime = time.Now() } + // Stamped here for the same reason as in rerank.go: the caller builds the + // request without a ModelConfig, this function has the one that loaded. + in.ModelIdentity = modelConfig.Model + res, err := depthModel.Depth(ctx, in) if appConfig.EnableTracing { diff --git a/core/backend/detection.go b/core/backend/detection.go index 1c92fefdfbad..4d0ddae33951 100644 --- a/core/backend/detection.go +++ b/core/backend/detection.go @@ -40,11 +40,12 @@ func Detection( } res, err := detectionModel.Detect(ctx, &proto.DetectOptions{ - Src: sourceFile, - Prompt: prompt, - Points: points, - Boxes: boxes, - Threshold: threshold, + ModelIdentity: modelConfig.Model, + Src: sourceFile, + Prompt: prompt, + Points: points, + Boxes: boxes, + Threshold: threshold, }) if appConfig.EnableTracing { diff --git a/core/backend/diarization.go b/core/backend/diarization.go index ba973d773403..241d1b20cf0f 100644 --- a/core/backend/diarization.go +++ b/core/backend/diarization.go @@ -19,19 +19,21 @@ import ( // don't act on. IncludeText only matters for backends that emit // per-segment transcripts as a by-product (e.g. vibevoice.cpp). type DiarizationRequest struct { - Audio string - Language string - NumSpeakers int32 - MinSpeakers int32 - MaxSpeakers int32 - ClusteringThreshold float32 - MinDurationOn float32 - MinDurationOff float32 - IncludeText bool + Audio string + Language string + NumSpeakers int32 + MinSpeakers int32 + MaxSpeakers int32 + ClusteringThreshold float32 + MinDurationOn float32 + MinDurationOff float32 + IncludeText bool } -func (r *DiarizationRequest) toProto(threads uint32) *proto.DiarizeRequest { +// modelIdentity: see the note on TranscriptionRequest.toProto. +func (r *DiarizationRequest) toProto(threads uint32, modelIdentity string) *proto.DiarizeRequest { return &proto.DiarizeRequest{ + ModelIdentity: modelIdentity, Dst: r.Audio, Threads: threads, Language: r.Language, @@ -74,7 +76,7 @@ func ModelDiarization(ctx context.Context, req DiarizationRequest, ml *model.Mod threads = uint32(*modelConfig.Threads) } - r, err := m.Diarize(ctx, req.toProto(threads)) + r, err := m.Diarize(ctx, req.toProto(threads, modelConfig.Model)) if err != nil { return nil, err } diff --git a/core/backend/face_analyze.go b/core/backend/face_analyze.go index bcd447243f26..c729506fe94b 100644 --- a/core/backend/face_analyze.go +++ b/core/backend/face_analyze.go @@ -37,9 +37,10 @@ func FaceAnalyze( } res, err := faceModel.FaceAnalyze(ctx, &proto.FaceAnalyzeRequest{ - Img: img, - Actions: actions, - AntiSpoofing: antiSpoofing, + ModelIdentity: modelConfig.Model, + Img: img, + Actions: actions, + AntiSpoofing: antiSpoofing, }) if appConfig.EnableTracing { diff --git a/core/backend/face_verify.go b/core/backend/face_verify.go index f99ed12d5a9c..0b51b694200d 100644 --- a/core/backend/face_verify.go +++ b/core/backend/face_verify.go @@ -37,10 +37,11 @@ func FaceVerify( } res, err := faceModel.FaceVerify(ctx, &proto.FaceVerifyRequest{ - Img1: img1, - Img2: img2, - Threshold: threshold, - AntiSpoofing: antiSpoofing, + ModelIdentity: modelConfig.Model, + Img1: img1, + Img2: img2, + Threshold: threshold, + AntiSpoofing: antiSpoofing, }) if appConfig.EnableTracing { diff --git a/core/backend/image.go b/core/backend/image.go index 7414d5d2b623..95f4dc122a85 100644 --- a/core/backend/image.go +++ b/core/backend/image.go @@ -29,6 +29,11 @@ func ImageGeneration(ctx context.Context, height, width, step, seed int, positiv _, err := inferenceModel.GenerateImage( ctx, &proto.GenerateImageRequest{ + // ModelIdentity is ModelConfig.Model — the SAME expression + // ModelOptions feeds to model.WithModel above, which the + // backend receives as ModelOptions.Model at LoadModel. Equal + // by construction, so the check cannot false-reject (#10952). + ModelIdentity: modelConfig.Model, Height: int32(height), Width: int32(width), Step: int32(step), diff --git a/core/backend/model_identity_modalities_test.go b/core/backend/model_identity_modalities_test.go new file mode 100644 index 000000000000..0087d84460ad --- /dev/null +++ b/core/backend/model_identity_modalities_test.go @@ -0,0 +1,346 @@ +package backend_test + +// Functional specs for the model identity every modality request must carry +// (#10952, extending #10970 beyond the four PredictOptions RPCs). +// +// The whole mechanism is only safe because the predict-time identity is the +// SAME expression as the load-time one: ModelOptions passes +// model.WithModel(c.Model), which becomes ModelOptions.Model at LoadModel, and +// each helper below builds its request in the same function from the same +// config value. If a helper ever sends ModelID()/Name or a resolved file path +// instead, every request to a correctly-routed backend gets rejected — so the +// config here deliberately gives Model, Name and ModelID() three different +// values, and these specs assert on the recorded wire value rather than on the +// construction site. + +import ( + "context" + "os" + "path/filepath" + + "github.com/mudler/LocalAI/core/backend" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/schema" + grpcPkg "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + "github.com/mudler/LocalAI/pkg/model" + "github.com/mudler/LocalAI/pkg/system" + ggrpc "google.golang.org/grpc" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +const ( + identityModelFile = "qwen/actual-weights.gguf" + identityModelName = "friendly-name" +) + +// recordingBackend captures the ModelIdentity each modality request arrives +// with. It embeds grpc.Backend so only the methods under test need declaring; +// anything else would nil-panic, which is the desired signal if a spec starts +// calling an unexpected RPC. +type recordingBackend struct { + grpcPkg.Backend + + seen map[string]string + + // SoundGenerationRequest carries both `model` (staged path) and + // `ModelIdentity`; the spec below asserts they are populated independently. + soundGenModelField string +} + +func newRecordingBackend() *recordingBackend { + return &recordingBackend{seen: map[string]string{}} +} + +// The loader probes liveness before handing the client back, so these two must +// answer even though no spec asserts on them. +func (r *recordingBackend) HealthCheck(context.Context) (bool, error) { return true, nil } +func (r *recordingBackend) IsBusy() bool { return false } + +func (r *recordingBackend) record(rpc string, in interface{ GetModelIdentity() string }) { + r.seen[rpc] = in.GetModelIdentity() +} + +func (r *recordingBackend) GenerateImage(_ context.Context, in *pb.GenerateImageRequest, _ ...ggrpc.CallOption) (*pb.Result, error) { + r.record("GenerateImage", in) + return &pb.Result{Success: true}, nil +} + +func (r *recordingBackend) SoundGeneration(_ context.Context, in *pb.SoundGenerationRequest, _ ...ggrpc.CallOption) (*pb.Result, error) { + r.record("SoundGeneration", in) + r.soundGenModelField = in.GetModel() + return &pb.Result{Success: true}, nil +} + +func (r *recordingBackend) GenerateVideo(_ context.Context, in *pb.GenerateVideoRequest, _ ...ggrpc.CallOption) (*pb.Result, error) { + r.record("GenerateVideo", in) + return &pb.Result{Success: true}, nil +} + +func (r *recordingBackend) Detect(_ context.Context, in *pb.DetectOptions, _ ...ggrpc.CallOption) (*pb.DetectResponse, error) { + r.record("Detect", in) + return &pb.DetectResponse{}, nil +} + +func (r *recordingBackend) Depth(_ context.Context, in *pb.DepthRequest, _ ...ggrpc.CallOption) (*pb.DepthResponse, error) { + r.record("Depth", in) + return &pb.DepthResponse{}, nil +} + +func (r *recordingBackend) VAD(_ context.Context, in *pb.VADRequest, _ ...ggrpc.CallOption) (*pb.VADResponse, error) { + r.record("VAD", in) + return &pb.VADResponse{}, nil +} + +func (r *recordingBackend) Diarize(_ context.Context, in *pb.DiarizeRequest, _ ...ggrpc.CallOption) (*pb.DiarizeResponse, error) { + r.record("Diarize", in) + return &pb.DiarizeResponse{}, nil +} + +func (r *recordingBackend) SoundDetection(_ context.Context, in *pb.SoundDetectionRequest, _ ...ggrpc.CallOption) (*pb.SoundDetectionResponse, error) { + r.record("SoundDetection", in) + return &pb.SoundDetectionResponse{}, nil +} + +func (r *recordingBackend) AudioTranscription(_ context.Context, in *pb.TranscriptRequest, _ ...ggrpc.CallOption) (*pb.TranscriptResult, error) { + r.record("AudioTranscription", in) + return &pb.TranscriptResult{}, nil +} + +func (r *recordingBackend) AudioTranscriptionStream(_ context.Context, in *pb.TranscriptRequest, _ func(*pb.TranscriptStreamResponse), _ ...ggrpc.CallOption) error { + r.record("AudioTranscriptionStream", in) + return nil +} + +func (r *recordingBackend) Rerank(_ context.Context, in *pb.RerankRequest, _ ...ggrpc.CallOption) (*pb.RerankResult, error) { + r.record("Rerank", in) + return &pb.RerankResult{Usage: &pb.Usage{}}, nil +} + +func (r *recordingBackend) Score(_ context.Context, in *pb.ScoreRequest, _ ...ggrpc.CallOption) (*pb.ScoreResponse, error) { + r.record("Score", in) + return &pb.ScoreResponse{}, nil +} + +func (r *recordingBackend) TokenClassify(_ context.Context, in *pb.TokenClassifyRequest, _ ...ggrpc.CallOption) (*pb.TokenClassifyResponse, error) { + r.record("TokenClassify", in) + return &pb.TokenClassifyResponse{}, nil +} + +func (r *recordingBackend) FaceVerify(_ context.Context, in *pb.FaceVerifyRequest, _ ...ggrpc.CallOption) (*pb.FaceVerifyResponse, error) { + r.record("FaceVerify", in) + return &pb.FaceVerifyResponse{}, nil +} + +func (r *recordingBackend) FaceAnalyze(_ context.Context, in *pb.FaceAnalyzeRequest, _ ...ggrpc.CallOption) (*pb.FaceAnalyzeResponse, error) { + r.record("FaceAnalyze", in) + return &pb.FaceAnalyzeResponse{}, nil +} + +func (r *recordingBackend) VoiceVerify(_ context.Context, in *pb.VoiceVerifyRequest, _ ...ggrpc.CallOption) (*pb.VoiceVerifyResponse, error) { + r.record("VoiceVerify", in) + return &pb.VoiceVerifyResponse{}, nil +} + +func (r *recordingBackend) VoiceAnalyze(_ context.Context, in *pb.VoiceAnalyzeRequest, _ ...ggrpc.CallOption) (*pb.VoiceAnalyzeResponse, error) { + r.record("VoiceAnalyze", in) + return &pb.VoiceAnalyzeResponse{}, nil +} + +func (r *recordingBackend) AudioTransform(_ context.Context, in *pb.AudioTransformRequest, _ ...ggrpc.CallOption) (*pb.AudioTransformResult, error) { + r.record("AudioTransform", in) + return &pb.AudioTransformResult{}, nil +} + +func (r *recordingBackend) VoiceEmbed(_ context.Context, in *pb.VoiceEmbedRequest, _ ...ggrpc.CallOption) (*pb.VoiceEmbedResponse, error) { + r.record("VoiceEmbed", in) + return &pb.VoiceEmbedResponse{Embedding: []float32{1}}, nil +} + +// newRecordingLoader wires a ModelLoader whose router hands back the recording +// backend, so every helper below reaches it through the real Load path. +func newRecordingLoader(rec *recordingBackend) *model.ModelLoader { + loader := model.NewModelLoader(&system.SystemState{}) + loader.SetModelRouter(func(_ context.Context, id string, _, _, _ string, _ *pb.ModelOptions, _ bool) (*model.Model, error) { + return model.NewModelWithClient(id, "test://recording", rec), nil + }) + return loader +} + +func identityModelCfg() config.ModelConfig { + threads := 1 + cfg := config.ModelConfig{ + Name: identityModelName, + Backend: "stub-backend", + Threads: &threads, + } + cfg.SetDefaults() + cfg.Name = identityModelName + cfg.Model = identityModelFile + return cfg +} + +var _ = Describe("modality requests carry the model identity", func() { + var ( + rec *recordingBackend + loader *model.ModelLoader + appCfg *config.ApplicationConfig + cfg config.ModelConfig + ctx context.Context + ) + + BeforeEach(func() { + rec = newRecordingBackend() + loader = newRecordingLoader(rec) + appCfg = config.NewApplicationConfig(config.WithSystemState(&system.SystemState{})) + cfg = identityModelCfg() + ctx = context.Background() + }) + + // Sanity: the three candidate values really are distinct here, so an + // assertion on ModelConfig.Model cannot pass by coincidence. + It("uses a config whose Model, Name and ModelID differ", func() { + Expect(cfg.Model).To(Equal(identityModelFile)) + Expect(cfg.Name).To(Equal(identityModelName)) + Expect(cfg.ModelID()).ToNot(Equal(cfg.Model)) + }) + + expectIdentity := func(rpc string) { + Expect(rec.seen).To(HaveKey(rpc), "%s never reached the backend", rpc) + Expect(rec.seen[rpc]).To(Equal(identityModelFile), + "%s must send ModelConfig.Model, the value LoadModel receives", rpc) + } + + It("ImageGeneration", func() { + fn, err := backend.ImageGeneration(ctx, 1, 1, 1, 1, "p", "n", "", "", loader, cfg, appCfg, nil) + Expect(err).ToNot(HaveOccurred()) + Expect(fn()).To(Succeed()) + expectIdentity("GenerateImage") + }) + + It("VideoGeneration", func() { + fn, err := backend.VideoGeneration(backend.VideoGenerationOptions{Prompt: "p"}, loader, cfg, appCfg) + Expect(err).ToNot(HaveOccurred()) + Expect(fn()).To(Succeed()) + expectIdentity("GenerateVideo") + }) + + It("Detection", func() { + _, err := backend.Detection(ctx, "src.png", "", nil, nil, 0, loader, appCfg, cfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("Detect") + }) + + It("Depth", func() { + _, err := backend.Depth(ctx, &pb.DepthRequest{Src: "src.png"}, loader, appCfg, cfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("Depth") + }) + + It("VAD", func() { + _, err := backend.VAD(&schema.VADRequest{Audio: []float32{0}}, ctx, loader, appCfg, cfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("VAD") + }) + + It("ModelDiarization", func() { + _, err := backend.ModelDiarization(ctx, backend.DiarizationRequest{Audio: "a.wav"}, loader, cfg, appCfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("Diarize") + }) + + It("ModelSoundDetection", func() { + _, err := backend.ModelSoundDetection(ctx, backend.SoundDetectionRequest{Audio: "a.wav"}, loader, cfg, appCfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("SoundDetection") + }) + + It("ModelTranscription", func() { + _, err := backend.ModelTranscription(ctx, "a.wav", "en", false, false, "", loader, cfg, appCfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("AudioTranscription") + }) + + It("ModelTranscriptionStream", func() { + err := backend.ModelTranscriptionStream(ctx, backend.TranscriptionRequest{Audio: "a.wav"}, loader, cfg, appCfg, func(backend.TranscriptionStreamChunk) {}) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("AudioTranscriptionStream") + }) + + It("Rerank", func() { + _, err := backend.Rerank(ctx, &pb.RerankRequest{Query: "q", Documents: []string{"d"}}, loader, appCfg, cfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("Rerank") + }) + + It("ModelScore", func() { + fn, err := backend.ModelScore("p", []string{"c"}, backend.ScoreOptions{}, loader, cfg, appCfg) + Expect(err).ToNot(HaveOccurred()) + _, err = fn(ctx) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("Score") + }) + + It("ModelTokenClassify", func() { + fn, err := backend.ModelTokenClassify("text", backend.TokenClassifyOptions{}, loader, cfg, appCfg) + Expect(err).ToNot(HaveOccurred()) + _, err = fn(ctx) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("TokenClassify") + }) + + It("FaceVerify", func() { + _, err := backend.FaceVerify(ctx, "img1", "img2", 0, false, loader, appCfg, cfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("FaceVerify") + }) + + It("FaceAnalyze", func() { + _, err := backend.FaceAnalyze(ctx, "img", nil, false, loader, appCfg, cfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("FaceAnalyze") + }) + + It("VoiceVerify", func() { + _, err := backend.VoiceVerify(ctx, "a1.wav", "a2.wav", 0, false, loader, appCfg, cfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("VoiceVerify") + }) + + It("VoiceAnalyze", func() { + _, err := backend.VoiceAnalyze(ctx, "a.wav", nil, loader, appCfg, cfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("VoiceAnalyze") + }) + + // SoundGeneration is the second RPC (with TTS) whose request already had a + // `model` field before this change. Identity must be a SEPARATE field: the + // distributed file stager rewrites `model` to a worker-local path, so + // comparing it would reject valid requests. Both must arrive populated. + It("SoundGeneration", func() { + appCfg.GeneratedContentDir = GinkgoT().TempDir() + _, _, err := backend.SoundGeneration(ctx, "hello", nil, nil, nil, nil, nil, nil, "", "", nil, "", "", "", nil, loader, appCfg, cfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("SoundGeneration") + Expect(rec.soundGenModelField).To(Equal(identityModelFile), + "the staged `model` field must still be sent, separately from the identity") + }) + + It("ModelAudioTransform", func() { + dir := GinkgoT().TempDir() + appCfg.GeneratedContentDir = dir + src := filepath.Join(dir, "in.wav") + Expect(os.WriteFile(src, []byte("RIFF"), 0o600)).To(Succeed()) + _, _, err := backend.ModelAudioTransform(ctx, src, "", backend.AudioTransformOptions{}, loader, appCfg, cfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("AudioTransform") + }) + + It("VoiceEmbed", func() { + _, err := backend.VoiceEmbed(ctx, "a.wav", loader, appCfg, cfg) + Expect(err).ToNot(HaveOccurred()) + expectIdentity("VoiceEmbed") + }) +}) diff --git a/core/backend/rerank.go b/core/backend/rerank.go index 0c0d15b226a1..94314276daba 100644 --- a/core/backend/rerank.go +++ b/core/backend/rerank.go @@ -77,6 +77,11 @@ func Rerank(ctx context.Context, request *proto.RerankRequest, loader *model.Mod startTime = time.Now() } + // Stamped here, not at the HTTP handler: this is the function that also + // builds ModelOptions from the same config, so the two values are equal by + // construction (#10952). + request.ModelIdentity = modelConfig.Model + res, err := rerankModel.Rerank(ctx, request) if appConfig.EnableTracing { diff --git a/core/backend/score.go b/core/backend/score.go index dce06213c715..7ce795f6098e 100644 --- a/core/backend/score.go +++ b/core/backend/score.go @@ -98,6 +98,7 @@ func ModelScore(prompt string, candidates []string, opts ScoreOptions, loader *m startTime = time.Now() } resp, err := b.Score(ctx, &pb.ScoreRequest{ + ModelIdentity: modelConfig.Model, Prompt: prompt, Candidates: candidates, IncludeTokenLogprobs: opts.IncludeTokenLogprobs, diff --git a/core/backend/sound_classification.go b/core/backend/sound_classification.go index 666c3232113a..581bdbc4ffea 100644 --- a/core/backend/sound_classification.go +++ b/core/backend/sound_classification.go @@ -22,11 +22,13 @@ type SoundDetectionRequest struct { Threshold float32 } -func (r *SoundDetectionRequest) toProto() *proto.SoundDetectionRequest { +// modelIdentity: see the note on TranscriptionRequest.toProto. +func (r *SoundDetectionRequest) toProto(modelIdentity string) *proto.SoundDetectionRequest { return &proto.SoundDetectionRequest{ - Src: r.Audio, - TopK: r.TopK, - Threshold: r.Threshold, + ModelIdentity: modelIdentity, + Src: r.Audio, + TopK: r.TopK, + Threshold: r.Threshold, } } @@ -54,7 +56,7 @@ func ModelSoundDetection(ctx context.Context, req SoundDetectionRequest, ml *mod return nil, err } - r, err := m.SoundDetection(ctx, req.toProto()) + r, err := m.SoundDetection(ctx, req.toProto(modelConfig.Model)) if err != nil { return nil, err } diff --git a/core/backend/soundgeneration.go b/core/backend/soundgeneration.go index 5a4ff8882c9b..e01e4b4a42e2 100644 --- a/core/backend/soundgeneration.go +++ b/core/backend/soundgeneration.go @@ -62,14 +62,19 @@ func SoundGeneration( } req := &proto.SoundGenerationRequest{ - Text: text, - Model: modelConfig.Model, - Dst: filePath, - Sample: doSample, - Duration: duration, - Temperature: temperature, - Src: sourceFile, - SrcDivisor: sourceDivisor, + Text: text, + // Model is rewritten to a worker-local path by the distributed file + // stager; ModelIdentity carries the untranslated value the backend + // compares against what it loaded (#10952). They start out equal here + // and must stay two separate fields. + ModelIdentity: modelConfig.Model, + Model: modelConfig.Model, + Dst: filePath, + Sample: doSample, + Duration: duration, + Temperature: temperature, + Src: sourceFile, + SrcDivisor: sourceDivisor, } if think != nil { req.Think = think diff --git a/core/backend/token_classify.go b/core/backend/token_classify.go index cb1e6b638c21..254b765cf669 100644 --- a/core/backend/token_classify.go +++ b/core/backend/token_classify.go @@ -87,8 +87,9 @@ func ModelTokenClassify(text string, opts TokenClassifyOptions, loader *model.Mo startTime = time.Now() } resp, err := inferenceModel.TokenClassify(ctx, &pb.TokenClassifyRequest{ - Text: text, - Threshold: opts.Threshold, + ModelIdentity: modelConfig.Model, + Text: text, + Threshold: opts.Threshold, }) entities := tokenClassifyResponseToEntities(resp) if appConfig.EnableTracing { diff --git a/core/backend/transcript.go b/core/backend/transcript.go index 211269160750..73826f3ae6eb 100644 --- a/core/backend/transcript.go +++ b/core/backend/transcript.go @@ -28,8 +28,14 @@ type TranscriptionRequest struct { TimestampGranularities []string } -func (r *TranscriptionRequest) toProto(threads uint32) *proto.TranscriptRequest { +// modelIdentity is ModelConfig.Model, the value the backend received as +// ModelOptions.Model at LoadModel, so it can reject a request that reached it +// through a stale distributed route (#10952). It is a parameter rather than a +// TranscriptionRequest field because the request is built by HTTP handlers that +// have no ModelConfig, while every caller of this method does. +func (r *TranscriptionRequest) toProto(threads uint32, modelIdentity string) *proto.TranscriptRequest { return &proto.TranscriptRequest{ + ModelIdentity: modelIdentity, Dst: r.Audio, Language: r.Language, Translate: r.Translate, @@ -85,7 +91,7 @@ func ModelTranscriptionWithOptions(ctx context.Context, req TranscriptionRequest audioSnippet = trace.AudioSnippet(req.Audio, appConfig.TracingMaxBodyBytes) } - r, err := transcriptionModel.AudioTranscription(ctx, req.toProto(uint32(*modelConfig.Threads))) + r, err := transcriptionModel.AudioTranscription(ctx, req.toProto(uint32(*modelConfig.Threads), modelConfig.Model)) if err != nil { if appConfig.EnableTracing { errData := map[string]any{ @@ -158,7 +164,7 @@ func ModelTranscriptionStream(ctx context.Context, req TranscriptionRequest, ml return err } - pbReq := req.toProto(uint32(*modelConfig.Threads)) + pbReq := req.toProto(uint32(*modelConfig.Threads), modelConfig.Model) pbReq.Stream = true return transcriptionModel.AudioTranscriptionStream(ctx, pbReq, func(chunk *proto.TranscriptStreamResponse) { diff --git a/core/backend/tts.go b/core/backend/tts.go index 2b49149ae46a..11ca0aabd45e 100644 --- a/core/backend/tts.go +++ b/core/backend/tts.go @@ -23,15 +23,22 @@ import ( // newTTSRequest assembles the gRPC TTSRequest from the per-request inputs. The // optional instructions string is only attached when non-empty so backends can // distinguish "no per-request instruction" (fall back to YAML) from an explicit -// empty one. params is forwarded as-is (nil when unset). -func newTTSRequest(text, modelPath, voice, dst, language, instructions string, params map[string]string) *proto.TTSRequest { +// empty one. params is forwarded as-is (nil when unset). modelIdentity is +// ModelConfig.Model - deliberately not modelPath, see the field comment below. +func newTTSRequest(text, modelPath, voice, dst, language, instructions string, params map[string]string, modelIdentity string) *proto.TTSRequest { req := &proto.TTSRequest{ - Text: text, - Model: modelPath, - Voice: voice, - Dst: dst, - Language: &language, - Params: params, + Text: text, + // Model is the path the backend synthesises from, and the distributed + // file stager rewrites it to a worker-local absolute path. Identity + // must therefore be a SEPARATE, untranslated field: comparing Model + // would reject valid requests in distributed mode, which is exactly + // the configuration this guards (#10952). + Model: modelPath, + ModelIdentity: modelIdentity, + Voice: voice, + Dst: dst, + Language: &language, + Params: params, } if instructions != "" { req.Instructions = &instructions @@ -95,7 +102,7 @@ func ModelTTS( startTime = time.Now() } - ttsRequest := newTTSRequest(text, modelPath, voice, filePath, language, instructions, params) + ttsRequest := newTTSRequest(text, modelPath, voice, filePath, language, instructions, params, modelConfig.Model) res, err := ttsModel.TTS(ctx, ttsRequest) @@ -197,7 +204,7 @@ func ModelTTSStream( snippetCapped := false // Streaming TTS writes to the HTTP response, not a file, so dst is empty. - ttsRequest := newTTSRequest(text, modelPath, voice, "", language, instructions, params) + ttsRequest := newTTSRequest(text, modelPath, voice, "", language, instructions, params, modelConfig.Model) err = ttsModel.TTSStream(ctx, ttsRequest, func(reply *proto.Reply) { // First message contains sample rate info diff --git a/core/backend/tts_test.go b/core/backend/tts_test.go index 8c3f35598c05..dc8daf15397b 100644 --- a/core/backend/tts_test.go +++ b/core/backend/tts_test.go @@ -13,7 +13,7 @@ import ( var _ = Describe("newTTSRequest", func() { It("attaches the instructions when a per-request value is set", func() { - req := newTTSRequest("hi", "/m", "alloy", "/out.wav", "en", "cheerful narrator", nil) + req := newTTSRequest("hi", "/m", "alloy", "/out.wav", "en", "cheerful narrator", nil, "m") Expect(req.Instructions).ToNot(BeNil()) Expect(req.GetInstructions()).To(Equal("cheerful narrator")) Expect(req.GetText()).To(Equal("hi")) @@ -23,20 +23,39 @@ var _ = Describe("newTTSRequest", func() { }) It("leaves instructions unset when empty so backends fall back to YAML", func() { - req := newTTSRequest("hi", "/m", "", "/out.wav", "", "", nil) + req := newTTSRequest("hi", "/m", "", "/out.wav", "", "", nil, "m") Expect(req.Instructions).To(BeNil()) Expect(req.GetInstructions()).To(Equal("")) }) It("forwards per-request params through to the backend", func() { params := map[string]string{"exaggeration": "0.7", "cfg_weight": "0.3"} - req := newTTSRequest("hi", "/m", "", "/out.wav", "", "", params) + req := newTTSRequest("hi", "/m", "", "/out.wav", "", "", params, "m") Expect(req.GetParams()).To(HaveKeyWithValue("exaggeration", "0.7")) Expect(req.GetParams()).To(HaveKeyWithValue("cfg_weight", "0.3")) }) It("leaves params nil when none are supplied", func() { - req := newTTSRequest("hi", "/m", "", "/out.wav", "", "", nil) + req := newTTSRequest("hi", "/m", "", "/out.wav", "", "", nil, "m") Expect(req.GetParams()).To(BeNil()) }) }) + +// TTSRequest carries TWO model strings and they are not interchangeable. +// `Model` is a path that FileStagingClient.TTS rewrites into a worker-local +// absolute path in distributed mode; `ModelIdentity` is the untranslated +// ModelConfig.Model the backend compares against what it loaded (#10952). +// Comparing `Model` instead would reject valid requests in exactly the +// configuration this guards, so these specs pin them as separate inputs. +var _ = Describe("newTTSRequest model identity", func() { + It("carries the untranslated identity alongside the staged model path", func() { + req := newTTSRequest("hi", "/worker/local/staged.onnx", "", "/out.wav", "", "", nil, "voices/piper-en.onnx") + Expect(req.GetModelIdentity()).To(Equal("voices/piper-en.onnx")) + Expect(req.GetModel()).To(Equal("/worker/local/staged.onnx")) + }) + + It("keeps the identity empty when the config names no model, so the backend skips", func() { + req := newTTSRequest("hi", "/m", "", "/out.wav", "", "", nil, "") + Expect(req.GetModelIdentity()).To(BeEmpty()) + }) +}) diff --git a/core/backend/vad.go b/core/backend/vad.go index 1c874e3382e5..c74055b76f85 100644 --- a/core/backend/vad.go +++ b/core/backend/vad.go @@ -25,7 +25,8 @@ func VAD(request *schema.VADRequest, } req := proto.VADRequest{ - Audio: request.Audio, + ModelIdentity: modelConfig.Model, + Audio: request.Audio, } resp, err := vadModel.VAD(ctx, &req) if err != nil { diff --git a/core/backend/video.go b/core/backend/video.go index 2702fefe9b8e..db56e363ca6c 100644 --- a/core/backend/video.go +++ b/core/backend/video.go @@ -41,6 +41,9 @@ func VideoGeneration(options VideoGenerationOptions, loader *model.ModelLoader, _, err := inferenceModel.GenerateVideo( appConfig.Context, &proto.GenerateVideoRequest{ + // See the ModelIdentity note in image.go: this is the value + // ModelOptions passed to LoadModel a few lines above. + ModelIdentity: modelConfig.Model, Height: options.Height, Width: options.Width, Prompt: options.Prompt, diff --git a/core/backend/voice_analyze.go b/core/backend/voice_analyze.go index 84e5547d5345..ddb20a6a9a0c 100644 --- a/core/backend/voice_analyze.go +++ b/core/backend/voice_analyze.go @@ -36,8 +36,9 @@ func VoiceAnalyze( } res, err := voiceModel.VoiceAnalyze(ctx, &proto.VoiceAnalyzeRequest{ - Audio: audio, - Actions: actions, + ModelIdentity: modelConfig.Model, + Audio: audio, + Actions: actions, }) if appConfig.EnableTracing { diff --git a/core/backend/voice_embed.go b/core/backend/voice_embed.go index 89df76a06a57..6608d4842ce2 100644 --- a/core/backend/voice_embed.go +++ b/core/backend/voice_embed.go @@ -39,7 +39,8 @@ func VoiceEmbed( } res, err := voiceModel.VoiceEmbed(ctx, &proto.VoiceEmbedRequest{ - Audio: audioPath, + ModelIdentity: modelConfig.Model, + Audio: audioPath, }) if appConfig.EnableTracing { diff --git a/core/backend/voice_verify.go b/core/backend/voice_verify.go index c5af1196bd21..e05380687305 100644 --- a/core/backend/voice_verify.go +++ b/core/backend/voice_verify.go @@ -37,10 +37,11 @@ func VoiceVerify( } res, err := voiceModel.VoiceVerify(ctx, &proto.VoiceVerifyRequest{ - Audio1: audio1, - Audio2: audio2, - Threshold: threshold, - AntiSpoofing: antiSpoofing, + ModelIdentity: modelConfig.Model, + Audio1: audio1, + Audio2: audio2, + Threshold: threshold, + AntiSpoofing: antiSpoofing, }) if appConfig.EnableTracing { diff --git a/pkg/grpc/model_identity_modalities_test.go b/pkg/grpc/model_identity_modalities_test.go new file mode 100644 index 000000000000..dd3ab6c0dc02 --- /dev/null +++ b/pkg/grpc/model_identity_modalities_test.go @@ -0,0 +1,226 @@ +package grpc + +import ( + "context" + + "github.com/mudler/LocalAI/pkg/grpc/base" + "github.com/mudler/LocalAI/pkg/grpc/grpcerrors" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +// modalityBackend answers every non-PredictOptions inference RPC successfully +// and counts the calls that reached it. `served` is therefore the signal for +// "the identity guard let this through" — the same role it plays in +// model_identity_test.go for the four PredictOptions RPCs. +type modalityBackend struct { + base.SingleThread + + loaded string + served int +} + +func (b *modalityBackend) Load(opts *pb.ModelOptions) error { + b.loaded = opts.Model + return nil +} + +func (b *modalityBackend) GenerateImage(*pb.GenerateImageRequest) error { b.served++; return nil } +func (b *modalityBackend) GenerateVideo(*pb.GenerateVideoRequest) error { b.served++; return nil } +func (b *modalityBackend) TTS(*pb.TTSRequest) error { b.served++; return nil } + +func (b *modalityBackend) TTSStream(_ *pb.TTSRequest, ch chan []byte) error { + b.served++ + ch <- []byte{0} + close(ch) + return nil +} + +func (b *modalityBackend) SoundGeneration(*pb.SoundGenerationRequest) error { b.served++; return nil } + +func (b *modalityBackend) AudioTranscription(context.Context, *pb.TranscriptRequest) (pb.TranscriptResult, error) { + b.served++ + return pb.TranscriptResult{Text: "ok"}, nil +} + +func (b *modalityBackend) AudioTranscriptionStream(_ context.Context, _ *pb.TranscriptRequest, ch chan *pb.TranscriptStreamResponse) error { + b.served++ + ch <- &pb.TranscriptStreamResponse{Delta: "ok"} + close(ch) + return nil +} + +func (b *modalityBackend) Detect(*pb.DetectOptions) (pb.DetectResponse, error) { + b.served++ + return pb.DetectResponse{}, nil +} + +func (b *modalityBackend) Depth(*pb.DepthRequest) (pb.DepthResponse, error) { + b.served++ + return pb.DepthResponse{}, nil +} + +func (b *modalityBackend) FaceVerify(*pb.FaceVerifyRequest) (pb.FaceVerifyResponse, error) { + b.served++ + return pb.FaceVerifyResponse{}, nil +} + +func (b *modalityBackend) FaceAnalyze(*pb.FaceAnalyzeRequest) (pb.FaceAnalyzeResponse, error) { + b.served++ + return pb.FaceAnalyzeResponse{}, nil +} + +func (b *modalityBackend) VoiceVerify(*pb.VoiceVerifyRequest) (pb.VoiceVerifyResponse, error) { + b.served++ + return pb.VoiceVerifyResponse{}, nil +} + +func (b *modalityBackend) VoiceAnalyze(*pb.VoiceAnalyzeRequest) (pb.VoiceAnalyzeResponse, error) { + b.served++ + return pb.VoiceAnalyzeResponse{}, nil +} + +func (b *modalityBackend) VoiceEmbed(*pb.VoiceEmbedRequest) (pb.VoiceEmbedResponse, error) { + b.served++ + return pb.VoiceEmbedResponse{}, nil +} + +func (b *modalityBackend) VAD(*pb.VADRequest) (pb.VADResponse, error) { + b.served++ + return pb.VADResponse{}, nil +} + +func (b *modalityBackend) Diarize(*pb.DiarizeRequest) (pb.DiarizeResponse, error) { + b.served++ + return pb.DiarizeResponse{}, nil +} + +func (b *modalityBackend) SoundDetection(context.Context, *pb.SoundDetectionRequest) (*pb.SoundDetectionResponse, error) { + b.served++ + return &pb.SoundDetectionResponse{}, nil +} + +func (b *modalityBackend) AudioTransform(*pb.AudioTransformRequest) (*pb.AudioTransformResult, error) { + b.served++ + return &pb.AudioTransformResult{}, nil +} + +var _ AIModel = (*modalityBackend)(nil) + +// callAllModalities exercises every remaining inference RPC that the generic Go +// server implements, sending `identity` in each request's ModelIdentity field. +// One call per RPC, so the returned error count is also the RPC count. +// +// Rerank, Score and TokenClassify are absent on purpose: the generic Go server +// does not implement them (they fall through to UnimplementedBackendServer), so +// only the C++ and Python backends can enforce them. +func callAllModalities(c Backend, identity string) map[string]error { + ctx := context.Background() + errs := map[string]error{} + + _, errs["GenerateImage"] = c.GenerateImage(ctx, &pb.GenerateImageRequest{ModelIdentity: identity}) + _, errs["GenerateVideo"] = c.GenerateVideo(ctx, &pb.GenerateVideoRequest{ModelIdentity: identity}) + _, errs["TTS"] = c.TTS(ctx, &pb.TTSRequest{ModelIdentity: identity}) + errs["TTSStream"] = c.TTSStream(ctx, &pb.TTSRequest{ModelIdentity: identity}, func(*pb.Reply) {}) + _, errs["SoundGeneration"] = c.SoundGeneration(ctx, &pb.SoundGenerationRequest{ModelIdentity: identity}) + _, errs["AudioTranscription"] = c.AudioTranscription(ctx, &pb.TranscriptRequest{ModelIdentity: identity}) + errs["AudioTranscriptionStream"] = c.AudioTranscriptionStream(ctx, &pb.TranscriptRequest{ModelIdentity: identity}, func(*pb.TranscriptStreamResponse) {}) + _, errs["Detect"] = c.Detect(ctx, &pb.DetectOptions{ModelIdentity: identity}) + _, errs["Depth"] = c.Depth(ctx, &pb.DepthRequest{ModelIdentity: identity}) + _, errs["FaceVerify"] = c.FaceVerify(ctx, &pb.FaceVerifyRequest{ModelIdentity: identity}) + _, errs["FaceAnalyze"] = c.FaceAnalyze(ctx, &pb.FaceAnalyzeRequest{ModelIdentity: identity}) + _, errs["VoiceVerify"] = c.VoiceVerify(ctx, &pb.VoiceVerifyRequest{ModelIdentity: identity}) + _, errs["VoiceAnalyze"] = c.VoiceAnalyze(ctx, &pb.VoiceAnalyzeRequest{ModelIdentity: identity}) + _, errs["VoiceEmbed"] = c.VoiceEmbed(ctx, &pb.VoiceEmbedRequest{ModelIdentity: identity}) + _, errs["VAD"] = c.VAD(ctx, &pb.VADRequest{ModelIdentity: identity}) + _, errs["Diarize"] = c.Diarize(ctx, &pb.DiarizeRequest{ModelIdentity: identity}) + _, errs["SoundDetection"] = c.SoundDetection(ctx, &pb.SoundDetectionRequest{ModelIdentity: identity}) + _, errs["AudioTransform"] = c.AudioTransform(ctx, &pb.AudioTransformRequest{ModelIdentity: identity}) + + return errs +} + +// #10970 protected only the four PredictOptions RPCs. Every other modality +// shares the exposure — the controller's cached row names a host:port, the port +// can be recycled by another model's backend, and a liveness-only probe cannot +// tell the difference — so each one must reject a request naming another model +// rather than answer from whatever it holds. +var _ = Describe("modality RPC model identity guard", func() { + newServed := func(addr, loadedModel string) (Backend, *modalityBackend) { + b := &modalityBackend{} + Provide(addr, b) + c := NewClient(addr, true, nil, false) + _, err := c.LoadModel(context.Background(), &pb.ModelOptions{Model: loadedModel}) + Expect(err).ToNot(HaveOccurred()) + Expect(b.loaded).To(Equal(loadedModel)) + return c, b + } + + It("rejects every modality RPC when the identity names another model", func() { + c, b := newServed("test://modality-mismatch", "a.gguf") + + for rpc, err := range callAllModalities(c, "b.gguf") { + Expect(err).To(HaveOccurred(), "%s must reject a wrong-model request", rpc) + Expect(grpcerrors.IsModelMismatch(err)).To(BeTrue(), "%s: want a mismatch error, got %v", rpc, err) + // The router reacts differently to the two signals, so a mismatch + // must never be mistaken for a not-loaded. + Expect(grpcerrors.IsModelNotLoaded(err)).To(BeFalse(), "%s", rpc) + } + Expect(b.served).To(Equal(0), "no request may reach the model on a mismatch") + }) + + It("serves when the identity matches the loaded model", func() { + c, b := newServed("test://modality-match", "a.gguf") + + errs := callAllModalities(c, "a.gguf") + for rpc, err := range errs { + Expect(err).ToNot(HaveOccurred(), "%s", rpc) + } + Expect(b.served).To(Equal(len(errs))) + }) + + // Compatibility, old controller -> new backend. Every existing deployment + // sends no identity, and tests/e2e-backends/backend_test.go drives real + // backends with bare request structs at many call sites. Tightening this + // breaks all of them, so it must fail here first. + It("serves when the request carries no identity", func() { + c, b := newServed("test://modality-empty-request", "a.gguf") + + errs := callAllModalities(c, "") + for rpc, err := range errs { + Expect(err).ToNot(HaveOccurred(), "%s", rpc) + } + Expect(b.served).To(Equal(len(errs))) + }) + + // The backend side of the same rule: a model loaded without an identity + // (an old controller did the load) cannot judge anything, so it must serve. + It("serves when the backend has no recorded identity", func() { + c, b := newServed("test://modality-empty-loaded", "") + + errs := callAllModalities(c, "b.gguf") + for rpc, err := range errs { + Expect(err).ToNot(HaveOccurred(), "%s", rpc) + } + Expect(b.served).To(Equal(len(errs))) + }) + + // AudioEncode/AudioDecode are deliberately NOT guarded: the opus codec + // backend they target is loaded with a literal model.WithModel("opus") and + // has no ModelConfig, so there is no value with the structural guarantee + // the rest of this mechanism depends on. Sending an identity that merely + // looks right would be convention, not construction — the exact + // false-rejection risk #10970 refused to take. Pinned here so a future + // "completeness" pass has to make that decision deliberately. + It("leaves the codec RPCs unguarded", func() { + c, _ := newServed("test://modality-codec", "a.gguf") + + // Whatever these answer, it must never be a model mismatch: no + // identity is carried, so nothing is compared. + _, err := c.AudioEncode(context.Background(), &pb.AudioEncodeRequest{}) + Expect(grpcerrors.IsModelMismatch(err)).To(BeFalse()) + _, err = c.AudioDecode(context.Background(), &pb.AudioDecodeRequest{}) + Expect(grpcerrors.IsModelMismatch(err)).To(BeFalse()) + }) +}) diff --git a/pkg/grpc/server.go b/pkg/grpc/server.go index 8bdd88f2fd95..9896a76dcd28 100644 --- a/pkg/grpc/server.go +++ b/pkg/grpc/server.go @@ -34,13 +34,20 @@ type server struct { pb.UnimplementedBackendServer llm AIModel - // identityMu guards loadedIdentity: LoadModel writes it, the - // PredictOptions RPCs read it, and nothing serialises those against each - // other (llm.Locking() is the model's own lock, and it is optional). + // identityMu guards loadedIdentity: LoadModel writes it, the inference + // RPCs read it, and nothing serialises those against each other + // (llm.Locking() is the model's own lock, and it is optional). identityMu sync.RWMutex loadedIdentity string } +// identifiedRequest is satisfied by every request message that carries a +// ModelIdentity field. protoc-gen-go generates GetModelIdentity with a +// nil-receiver guard, so a typed-nil request is safe to pass here. +type identifiedRequest interface { + GetModelIdentity() string +} + // checkModelIdentity reports an error when the request names a model other // than the one this process loaded. It is the point-of-use half of the fix for // #10952: in distributed mode a worker can recycle a stopped backend's gRPC @@ -53,18 +60,23 @@ type server struct { // requests; the loaded side is empty when such a controller performed the // load. Neither can judge the other, and a false rejection is far worse than // the miss. -func (s *server) checkModelIdentity(in *pb.PredictOptions) error { - if in == nil || in.ModelIdentity == "" { +// +// One guard covers every modality because they all share one exposure: the +// route is cached by address, the address can be recycled, and only the +// backend can tell. Which RPCs call it is the whole enforcement surface — see +// pkg/grpc/model_identity_modalities_test.go, which drives all of them. +func (s *server) checkModelIdentity(in identifiedRequest) error { + if in == nil || in.GetModelIdentity() == "" { return nil } s.identityMu.RLock() loaded := s.loadedIdentity s.identityMu.RUnlock() - if loaded == "" || loaded == in.ModelIdentity { + if loaded == "" || loaded == in.GetModelIdentity() { return nil } - return grpcerrors.ModelMismatch("grpc-server", loaded, in.ModelIdentity) + return grpcerrors.ModelMismatch("grpc-server", loaded, in.GetModelIdentity()) } func (s *server) Health(ctx context.Context, in *pb.HealthMessage) (*pb.Reply, error) { @@ -124,6 +136,9 @@ func (s *server) Predict(ctx context.Context, in *pb.PredictOptions) (*pb.Reply, } func (s *server) GenerateImage(ctx context.Context, in *pb.GenerateImageRequest) (*pb.Result, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -136,6 +151,9 @@ func (s *server) GenerateImage(ctx context.Context, in *pb.GenerateImageRequest) } func (s *server) GenerateVideo(ctx context.Context, in *pb.GenerateVideoRequest) (*pb.Result, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -148,6 +166,9 @@ func (s *server) GenerateVideo(ctx context.Context, in *pb.GenerateVideoRequest) } func (s *server) TTS(ctx context.Context, in *pb.TTSRequest) (*pb.Result, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -160,6 +181,9 @@ func (s *server) TTS(ctx context.Context, in *pb.TTSRequest) (*pb.Result, error) } func (s *server) TTSStream(in *pb.TTSRequest, stream pb.Backend_TTSStreamServer) error { + if err := s.checkModelIdentity(in); err != nil { + return err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -181,6 +205,9 @@ func (s *server) TTSStream(in *pb.TTSRequest, stream pb.Backend_TTSStreamServer) } func (s *server) SoundGeneration(ctx context.Context, in *pb.SoundGenerationRequest) (*pb.Result, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -193,6 +220,9 @@ func (s *server) SoundGeneration(ctx context.Context, in *pb.SoundGenerationRequ } func (s *server) Detect(ctx context.Context, in *pb.DetectOptions) (*pb.DetectResponse, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -205,6 +235,9 @@ func (s *server) Detect(ctx context.Context, in *pb.DetectOptions) (*pb.DetectRe } func (s *server) Depth(ctx context.Context, in *pb.DepthRequest) (*pb.DepthResponse, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -217,6 +250,9 @@ func (s *server) Depth(ctx context.Context, in *pb.DepthRequest) (*pb.DepthRespo } func (s *server) FaceVerify(ctx context.Context, in *pb.FaceVerifyRequest) (*pb.FaceVerifyResponse, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -229,6 +265,9 @@ func (s *server) FaceVerify(ctx context.Context, in *pb.FaceVerifyRequest) (*pb. } func (s *server) FaceAnalyze(ctx context.Context, in *pb.FaceAnalyzeRequest) (*pb.FaceAnalyzeResponse, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -241,6 +280,9 @@ func (s *server) FaceAnalyze(ctx context.Context, in *pb.FaceAnalyzeRequest) (*p } func (s *server) VoiceVerify(ctx context.Context, in *pb.VoiceVerifyRequest) (*pb.VoiceVerifyResponse, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -253,6 +295,9 @@ func (s *server) VoiceVerify(ctx context.Context, in *pb.VoiceVerifyRequest) (*p } func (s *server) VoiceAnalyze(ctx context.Context, in *pb.VoiceAnalyzeRequest) (*pb.VoiceAnalyzeResponse, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -265,6 +310,9 @@ func (s *server) VoiceAnalyze(ctx context.Context, in *pb.VoiceAnalyzeRequest) ( } func (s *server) VoiceEmbed(ctx context.Context, in *pb.VoiceEmbedRequest) (*pb.VoiceEmbedResponse, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -277,6 +325,9 @@ func (s *server) VoiceEmbed(ctx context.Context, in *pb.VoiceEmbedRequest) (*pb. } func (s *server) AudioTranscription(ctx context.Context, in *pb.TranscriptRequest) (*pb.TranscriptResult, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -319,6 +370,9 @@ func (s *server) AudioTranscription(ctx context.Context, in *pb.TranscriptReques } func (s *server) AudioTranscriptionStream(in *pb.TranscriptRequest, stream pb.Backend_AudioTranscriptionStreamServer) error { + if err := s.checkModelIdentity(in); err != nil { + return err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -536,6 +590,9 @@ func (s *server) StoresFind(ctx context.Context, in *pb.StoresFindOptions) (*pb. } func (s *server) VAD(ctx context.Context, in *pb.VADRequest) (*pb.VADResponse, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -548,6 +605,9 @@ func (s *server) VAD(ctx context.Context, in *pb.VADRequest) (*pb.VADResponse, e } func (s *server) Diarize(ctx context.Context, in *pb.DiarizeRequest) (*pb.DiarizeResponse, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -560,6 +620,9 @@ func (s *server) Diarize(ctx context.Context, in *pb.DiarizeRequest) (*pb.Diariz } func (s *server) SoundDetection(ctx context.Context, in *pb.SoundDetectionRequest) (*pb.SoundDetectionResponse, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() @@ -592,6 +655,9 @@ func (s *server) AudioDecode(ctx context.Context, in *pb.AudioDecodeRequest) (*p } func (s *server) AudioTransform(ctx context.Context, in *pb.AudioTransformRequest) (*pb.AudioTransformResult, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } if s.llm.Locking() { s.llm.Lock() defer s.llm.Unlock() diff --git a/tests/e2e/mock-backend/main.go b/tests/e2e/mock-backend/main.go index 32d52cc126e7..dbe405e10a36 100644 --- a/tests/e2e/mock-backend/main.go +++ b/tests/e2e/mock-backend/main.go @@ -59,16 +59,19 @@ func snapshotLoadParams() *pb.ModelOptions { // backend/python/common/model_identity.py, backend/cpp/*/grpc-server.cpp) so // the distributed e2e suite exercises the #10952 fix rather than only the // unit tests. Empty on either side means skip, which is why every existing -// spec that sends a bare PredictOptions keeps working. -func checkModelIdentity(in *pb.PredictOptions) error { - if in == nil || in.ModelIdentity == "" { +// spec that sends a bare request struct keeps working. +// +// It takes an interface rather than *pb.PredictOptions because every modality +// request message now carries the field, and the rule must not drift per RPC. +func checkModelIdentity(in interface{ GetModelIdentity() string }) error { + if in == nil || in.GetModelIdentity() == "" { return nil } opts := snapshotLoadParams() - if opts == nil || opts.Model == "" || opts.Model == in.ModelIdentity { + if opts == nil || opts.Model == "" || opts.Model == in.GetModelIdentity() { return nil } - return grpcerrors.ModelMismatch("mock-backend", opts.Model, in.ModelIdentity) + return grpcerrors.ModelMismatch("mock-backend", opts.Model, in.GetModelIdentity()) } // promptHasToolResults checks if the prompt contains evidence of prior tool @@ -427,6 +430,9 @@ func (m *MockBackend) Embedding(ctx context.Context, in *pb.PredictOptions) (*pb } func (m *MockBackend) GenerateImage(ctx context.Context, in *pb.GenerateImageRequest) (*pb.Result, error) { + if err := checkModelIdentity(in); err != nil { + return nil, err + } xlog.Debug("GenerateImage called", "prompt", in.PositivePrompt) return &pb.Result{ Message: "Image generated successfully (mocked)", @@ -435,6 +441,9 @@ func (m *MockBackend) GenerateImage(ctx context.Context, in *pb.GenerateImageReq } func (m *MockBackend) GenerateVideo(ctx context.Context, in *pb.GenerateVideoRequest) (*pb.Result, error) { + if err := checkModelIdentity(in); err != nil { + return nil, err + } xlog.Debug("GenerateVideo called", "prompt", in.Prompt) return &pb.Result{ Message: "Video generated successfully (mocked)", @@ -443,6 +452,9 @@ func (m *MockBackend) GenerateVideo(ctx context.Context, in *pb.GenerateVideoReq } func (m *MockBackend) TTS(ctx context.Context, in *pb.TTSRequest) (*pb.Result, error) { + if err := checkModelIdentity(in); err != nil { + return nil, err + } xlog.Debug("TTS called", "text", in.Text) dst := in.GetDst() if dst != "" { @@ -460,6 +472,9 @@ func (m *MockBackend) TTS(ctx context.Context, in *pb.TTSRequest) (*pb.Result, e } func (m *MockBackend) TTSStream(in *pb.TTSRequest, stream pb.Backend_TTSStreamServer) error { + if err := checkModelIdentity(in); err != nil { + return err + } xlog.Debug("TTSStream called", "text", in.Text) // Stream mock audio chunks (simplified - just send a few bytes) chunks := [][]byte{ @@ -476,6 +491,9 @@ func (m *MockBackend) TTSStream(in *pb.TTSRequest, stream pb.Backend_TTSStreamSe } func (m *MockBackend) SoundGeneration(ctx context.Context, in *pb.SoundGenerationRequest) (*pb.Result, error) { + if err := checkModelIdentity(in); err != nil { + return nil, err + } xlog.Debug("SoundGeneration called", "text", in.Text, "caption", in.GetCaption(), @@ -555,6 +573,9 @@ func writeMinimalWAV(path string) error { } func (m *MockBackend) AudioTranscription(ctx context.Context, in *pb.TranscriptRequest) (*pb.TranscriptResult, error) { + if err := checkModelIdentity(in); err != nil { + return nil, err + } dst := in.GetDst() wavSR := 0 dataLen := 0 @@ -629,6 +650,9 @@ func (m *MockBackend) TokenizeString(ctx context.Context, in *pb.PredictOptions) // scores equally, softmax is uniform, and (with a sensible activation // threshold) the router falls back. func (m *MockBackend) Score(ctx context.Context, in *pb.ScoreRequest) (*pb.ScoreResponse, error) { + if err := checkModelIdentity(in); err != nil { + return nil, err + } xlog.Debug("Score called", "candidates", len(in.Candidates)) hint := extractRouteHint(in.Prompt) out := &pb.ScoreResponse{Candidates: make([]*pb.CandidateScore, len(in.Candidates))} @@ -702,6 +726,9 @@ func (m *MockBackend) Status(ctx context.Context, in *pb.HealthMessage) (*pb.Sta } func (m *MockBackend) Detect(ctx context.Context, in *pb.DetectOptions) (*pb.DetectResponse, error) { + if err := checkModelIdentity(in); err != nil { + return nil, err + } xlog.Debug("Detect called", "src", in.Src) return &pb.DetectResponse{ Detections: []*pb.Detection{ @@ -770,6 +797,9 @@ func (m *MockBackend) StoresFind(ctx context.Context, in *pb.StoresFindOptions) } func (m *MockBackend) Rerank(ctx context.Context, in *pb.RerankRequest) (*pb.RerankResult, error) { + if err := checkModelIdentity(in); err != nil { + return nil, err + } xlog.Debug("Rerank called", "query", in.Query, "documents", len(in.Documents)) // Return mock reranking results results := make([]*pb.DocumentResult, len(in.Documents)) @@ -801,6 +831,9 @@ func (m *MockBackend) GetMetrics(ctx context.Context, in *pb.MetricsRequest) (*p } func (m *MockBackend) VAD(ctx context.Context, in *pb.VADRequest) (*pb.VADResponse, error) { + if err := checkModelIdentity(in); err != nil { + return nil, err + } // Compute RMS of the received float32 audio to decide whether speech is present. var sumSq float64 for _, s := range in.Audio { @@ -836,6 +869,9 @@ func (m *MockBackend) VAD(ctx context.Context, in *pb.VADRequest) (*pb.VADRespon // should reflect two segments (1.0s + 1.5s = 2.5s), and IncludeText must // gate the per-segment Text field. func (m *MockBackend) Diarize(ctx context.Context, in *pb.DiarizeRequest) (*pb.DiarizeResponse, error) { + if err := checkModelIdentity(in); err != nil { + return nil, err + } xlog.Debug("Diarize called", "dst", in.Dst, "num_speakers", in.NumSpeakers, @@ -933,6 +969,9 @@ func voiceEmbedFromWAV(path string) []float32 { // VoiceEmbed returns a deterministic 2-d speaker embedding for the audio clip. // See voiceEmbedFromWAV for the (test-only) DC-offset discrimination scheme. func (m *MockBackend) VoiceEmbed(ctx context.Context, in *pb.VoiceEmbedRequest) (*pb.VoiceEmbedResponse, error) { + if err := checkModelIdentity(in); err != nil { + return nil, err + } emb := voiceEmbedFromWAV(in.GetAudio()) xlog.Debug("VoiceEmbed called", "audio", in.GetAudio(), "embedding", emb) if len(emb) == 0 { @@ -943,6 +982,9 @@ func (m *MockBackend) VoiceEmbed(ctx context.Context, in *pb.VoiceEmbedRequest) // VoiceVerify compares two clips by cosine distance over their mock embeddings. func (m *MockBackend) VoiceVerify(ctx context.Context, in *pb.VoiceVerifyRequest) (*pb.VoiceVerifyResponse, error) { + if err := checkModelIdentity(in); err != nil { + return nil, err + } a := voiceEmbedFromWAV(in.GetAudio1()) b := voiceEmbedFromWAV(in.GetAudio2()) dist := float32(1)