fix(distributed): reject wrong-model requests on the remaining modalities (#10990)

#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.


Assisted-by: Claude Code:claude-opus-4-8 [Read] [Edit] [Bash]

Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
Co-authored-by: Ettore Di Giacinto <mudler@localai.io>
This commit is contained in:
mudler's LocalAI [bot]
2026-07-20 21:58:19 +02:00
committed by GitHub
parent a784cf669f
commit 1cd7d63c7b
29 changed files with 1104 additions and 85 deletions

View File

@@ -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)