From 9ff4af5734a28f1cbd31d95d5dfc163cd309c9c9 Mon Sep 17 00:00:00 2001 From: "ilblackdragon@gmail.com" Date: Mon, 23 Mar 2026 23:29:26 -0700 Subject: [PATCH] fix: heartbeat hygiene, /model multi-tenant guard, RigAdapter model override Three follow-up fixes for multi-tenant isolation: 1. Multi-user heartbeat now runs memory hygiene per user before each heartbeat check, matching single-user heartbeat behavior. 2. /model command in multi-tenant mode only persists to per-user settings (selected_model) without calling set_model() on the shared LlmProvider. The per-request model_override in the dispatcher reads from the same setting. Added multi_tenant flag to AgentConfig (auto-detected from GATEWAY_USER_TOKENS). 3. RigAdapter now supports per-request model overrides by injecting the model name into rig-core's additional_params. OpenAI/Anthropic/Ollama API servers use last-key-wins for duplicate JSON keys, so the override takes effect via serde's flatten serialization order. Co-Authored-By: Claude Opus 4.6 (1M context) --- src/agent/commands.rs | 37 ++++++++++++++++++++---------- src/agent/dispatcher.rs | 3 +++ src/agent/heartbeat.rs | 17 ++++++++++++++ src/agent/scheduler.rs | 1 + src/config/agent.rs | 10 +++++++- src/llm/rig_adapter.rs | 51 +++++++++++++++++++++++++---------------- 6 files changed, 86 insertions(+), 33 deletions(-) diff --git a/src/agent/commands.rs b/src/agent/commands.rs index c957a40e..4e5b681c 100644 --- a/src/agent/commands.rs +++ b/src/agent/commands.rs @@ -663,19 +663,32 @@ impl Agent { } } - match self.llm().set_model(requested) { - Ok(()) => { - // Persist the model choice so it survives restarts. - self.persist_selected_model(requested).await; - Ok(SubmissionResult::response(format!( - "Switched model to: {}", - requested - ))) + if self.config.multi_tenant { + // Multi-tenant: only persist to per-user settings. + // Do NOT call set_model() on the shared provider — that + // would change the default for all users. The per-request + // model_override in the dispatcher reads from the same + // "selected_model" setting and applies it per-user. + self.persist_selected_model(requested).await; + Ok(SubmissionResult::response(format!( + "Model preference set to: {} (per-user)", + requested + ))) + } else { + match self.llm().set_model(requested) { + Ok(()) => { + // Persist the model choice so it survives restarts. + self.persist_selected_model(requested).await; + Ok(SubmissionResult::response(format!( + "Switched model to: {}", + requested + ))) + } + Err(e) => Ok(SubmissionResult::error(format!( + "Failed to switch model: {}", + e + ))), } - Err(e) => Ok(SubmissionResult::error(format!( - "Failed to switch model: {}", - e - ))), } } } diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index 4bd0a819..d6e05c9b 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -1277,6 +1277,7 @@ mod tests { auto_approve_tools: false, default_timezone: "UTC".to_string(), max_tokens_per_job: 0, + multi_tenant: false, }, deps, Arc::new(ChannelManager::new()), @@ -2146,6 +2147,7 @@ mod tests { auto_approve_tools: true, default_timezone: "UTC".to_string(), max_tokens_per_job: 0, + multi_tenant: false, }, deps, Arc::new(ChannelManager::new()), @@ -2268,6 +2270,7 @@ mod tests { auto_approve_tools: true, default_timezone: "UTC".to_string(), max_tokens_per_job: 0, + multi_tenant: false, }, deps, Arc::new(ChannelManager::new()), diff --git a/src/agent/heartbeat.rs b/src/agent/heartbeat.rs index 443cb434..d58ed4db 100644 --- a/src/agent/heartbeat.rs +++ b/src/agent/heartbeat.rs @@ -572,6 +572,23 @@ pub fn spawn_multi_user_heartbeat( for user_id in &user_ids { let workspace = Arc::new(Workspace::new_with_db(user_id, store.clone())); + // Run memory hygiene per user (same as single-user heartbeat). + let hygiene_ws = Arc::clone(&workspace); + let hygiene_cfg = hygiene_config.clone(); + let hygiene_user = user_id.clone(); + tokio::spawn(async move { + let report = + crate::workspace::hygiene::run_if_due(&hygiene_ws, &hygiene_cfg).await; + if report.had_work() { + tracing::info!( + user_id = hygiene_user, + daily_logs_deleted = report.daily_logs_deleted, + conversation_docs_deleted = report.conversation_docs_deleted, + "multi-user heartbeat: memory hygiene deleted stale documents" + ); + } + }); + let mut runner = HeartbeatRunner::new( config.clone(), hygiene_config.clone(), diff --git a/src/agent/scheduler.rs b/src/agent/scheduler.rs index db5e049a..a082fe23 100644 --- a/src/agent/scheduler.rs +++ b/src/agent/scheduler.rs @@ -785,6 +785,7 @@ mod tests { auto_approve_tools: true, default_timezone: "UTC".to_string(), max_tokens_per_job, + multi_tenant: false, }; let cm = Arc::new(ContextManager::new(5)); let llm: Arc = Arc::new(StubLlm); diff --git a/src/config/agent.rs b/src/config/agent.rs index 6a724c25..e91b9140 100644 --- a/src/config/agent.rs +++ b/src/config/agent.rs @@ -1,6 +1,6 @@ use std::time::Duration; -use crate::config::helpers::{parse_bool_env, parse_option_env, parse_optional_env}; +use crate::config::helpers::{optional_env, parse_bool_env, parse_option_env, parse_optional_env}; use crate::error::ConfigError; use crate::settings::Settings; @@ -33,6 +33,9 @@ pub struct AgentConfig { pub default_timezone: String, /// Maximum tokens per job (0 = unlimited). pub max_tokens_per_job: u64, + /// Whether the deployment is multi-tenant (multiple users sharing one + /// instance). Auto-detected from GATEWAY_USER_TOKENS presence. + pub multi_tenant: bool, } impl AgentConfig { @@ -56,6 +59,7 @@ impl AgentConfig { auto_approve_tools: true, default_timezone: "UTC".to_string(), max_tokens_per_job: 0, + multi_tenant: false, } } @@ -116,6 +120,10 @@ impl AgentConfig { "AGENT_MAX_TOKENS_PER_JOB", settings.agent.max_tokens_per_job, )?, + multi_tenant: parse_bool_env( + "MULTI_TENANT", + optional_env("GATEWAY_USER_TOKENS")?.is_some(), + )?, }) } } diff --git a/src/llm/rig_adapter.rs b/src/llm/rig_adapter.rs index a9030929..07aec57f 100644 --- a/src/llm/rig_adapter.rs +++ b/src/llm/rig_adapter.rs @@ -597,6 +597,29 @@ fn build_rig_request( }) } +/// Inject a per-request model override into the rig request's `additional_params`. +/// +/// Rig-core bakes the model name at construction time. For OpenAI, Anthropic, and +/// Ollama, the `model` field in the request body determines which model serves the +/// request. Rig-core's `#[serde(flatten)]` on `additional_params` emits these fields +/// AFTER the struct's own `model` field. Most API servers (Python, Go) use +/// last-key-wins when deserializing duplicate JSON keys, so the override takes effect. +fn inject_model_override(rig_req: &mut RigRequest, model_override: Option<&str>) { + let Some(model) = model_override else { + return; + }; + match rig_req.additional_params { + Some(ref mut params) => { + if let Some(obj) = params.as_object_mut() { + obj.insert("model".to_string(), serde_json::json!(model)); + } + } + None => { + rig_req.additional_params = Some(serde_json::json!({ "model": model })); + } + } +} + #[async_trait] impl LlmProvider for RigAdapter where @@ -631,15 +654,7 @@ where &self, mut request: CompletionRequest, ) -> Result { - if let Some(requested_model) = request.model.as_deref() - && requested_model != self.model_name.as_str() - { - tracing::warn!( - requested_model = requested_model, - active_model = %self.model_name, - "Per-request model override is not supported for this provider; using configured model" - ); - } + let model_override = request.model.take(); self.strip_unsupported_completion_params(&mut request); @@ -647,7 +662,7 @@ where crate::llm::provider::sanitize_tool_messages(&mut messages); let (preamble, history) = convert_messages(&messages); - let rig_req = build_rig_request( + let mut rig_req = build_rig_request( preamble, history, Vec::new(), @@ -657,6 +672,8 @@ where self.cache_retention, )?; + inject_model_override(&mut rig_req, model_override.as_deref()); + let response = self.model .completion(rig_req) @@ -694,15 +711,7 @@ where &self, mut request: ToolCompletionRequest, ) -> Result { - if let Some(requested_model) = request.model.as_deref() - && requested_model != self.model_name.as_str() - { - tracing::warn!( - requested_model = requested_model, - active_model = %self.model_name, - "Per-request model override is not supported for this provider; using configured model" - ); - } + let model_override = request.model.take(); self.strip_unsupported_tool_params(&mut request); @@ -715,7 +724,7 @@ where let tools = convert_tools(&request.tools); let tool_choice = convert_tool_choice(request.tool_choice.as_deref()); - let rig_req = build_rig_request( + let mut rig_req = build_rig_request( preamble, history, tools, @@ -725,6 +734,8 @@ where self.cache_retention, )?; + inject_model_override(&mut rig_req, model_override.as_deref()); + let response = self.model .completion(rig_req)