diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index 00442c41..5713d599 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -398,16 +398,21 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { }; // Record cost and track token usage (global + per-user). - // When a model override is active, use the override name for attribution - // and let CostGuard look up pricing via costs::model_cost() instead of - // using the default provider's cost_per_token (which reflects the wrong model). - let (model_name, cost_per_token) = if let Some(ref ovr) = reason_ctx.model_override { - (ovr.clone(), None) + // Use the provider's effective_model_name so cost attribution matches + // the model that actually served the request. When the override is + // honoured (e.g. NearAI), this returns the override name; when the + // provider ignores overrides (e.g. Rig-based), it returns the active + // model, keeping attribution accurate in both cases. + let model_name = self + .agent + .llm() + .effective_model_name(reason_ctx.model_override.as_deref()); + let cost_per_token = if reason_ctx.model_override.is_some() { + // Override may use different pricing; let CostGuard fall back to + // costs::model_cost() for the effective model. + None } else { - ( - self.agent.llm().active_model_name(), - Some(self.agent.llm().cost_per_token()), - ) + Some(self.agent.llm().cost_per_token()) }; let read_discount = self.agent.llm().cache_read_discount(); let write_multiplier = self.agent.llm().cache_write_multiplier(); diff --git a/src/agent/heartbeat.rs b/src/agent/heartbeat.rs index f7a8f869..18fafc1d 100644 --- a/src/agent/heartbeat.rs +++ b/src/agent/heartbeat.rs @@ -400,7 +400,7 @@ impl HeartbeatRunner { } /// Send a notification about heartbeat findings. - pub(crate) async fn send_notification(&self, message: &str) { + async fn send_notification(&self, message: &str) { let Some(ref tx) = self.response_tx else { tracing::debug!("No response channel configured for heartbeat notifications"); return; @@ -512,8 +512,8 @@ pub fn spawn_heartbeat( }) } -/// Spawn a multi-user heartbeat runner that cycles through all users that -/// own routines (enabled or not). Each tick, it queries the DB for distinct +/// Spawn a multi-user heartbeat runner that cycles through all users who +/// have routines (enabled or not). Each tick, it queries the DB for distinct /// user_ids, creates a per-user workspace, and runs a heartbeat check for /// each user concurrently. Per-user failure counts are tracked independently. pub fn spawn_multi_user_heartbeat( @@ -574,8 +574,10 @@ pub fn spawn_multi_user_heartbeat( } }; - // Run user heartbeats concurrently so one slow LLM call doesn't - // block others. Cap concurrency to avoid flooding the LLM provider. + // Run user heartbeats (and hygiene) concurrently so one slow LLM + // call doesn't block others. Cap concurrency to avoid flooding the + // LLM provider. Hygiene runs inside the same JoinSet so it is + // tracked and bounded by the same concurrency cap. const MAX_CONCURRENT_HEARTBEATS: usize = 8; let mut join_set = tokio::task::JoinSet::new(); @@ -588,23 +590,6 @@ pub fn spawn_multi_user_heartbeat( let workspace = Arc::new(Workspace::new_with_db(user_id, Arc::clone(store.db()))); - // 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" - ); - } - }); - // Drain completed tasks to stay within the concurrency cap. while join_set.len() >= MAX_CONCURRENT_HEARTBEATS { if let Some(join_result) = join_set.join_next().await { @@ -613,13 +598,30 @@ pub fn spawn_multi_user_heartbeat( } let uid = user_id.clone(); - let cfg = config.clone(); + // In multi-tenant mode, clear notify_user_id so that + // HeartbeatRunner::send_notification falls back to + // workspace.user_id() — each user's heartbeat should persist + // and notify that user, not the shared config target. + let mut cfg = config.clone(); + cfg.notify_user_id = None; let hyg = hygiene_config.clone(); let llm_clone = llm.clone(); let tx = response_tx.clone(); let admin = store.clone(); join_set.spawn(async move { + // Run memory hygiene per user (same as single-user heartbeat) + // inside the tracked task so concurrency is bounded. + let report = crate::workspace::hygiene::run_if_due(&workspace, &hyg).await; + if report.had_work() { + tracing::info!( + user_id = uid, + 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(cfg, hyg, workspace, llm_clone); if let Some(tx) = tx { runner = runner.with_response_channel(tx); diff --git a/src/llm/rig_adapter.rs b/src/llm/rig_adapter.rs index 038236fd..bfc6c567 100644 --- a/src/llm/rig_adapter.rs +++ b/src/llm/rig_adapter.rs @@ -601,11 +601,14 @@ 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 inside each provider's -/// `CompletionModel` implementation. The actual HTTP request body includes a -/// `model` field set by the provider. Rig-core's `#[serde(flatten)]` on -/// `additional_params` emits these fields AFTER the provider's own fields. -/// Most API servers (Python, Go) use last-key-wins when deserializing -/// duplicate JSON keys, so the injected `model` value takes effect. +/// `CompletionModel` implementation. This helper inserts a top-level `"model"` +/// key into `additional_params`, which rig-core flattens into the provider's +/// request payload via `#[serde(flatten)]`. +/// +/// Whether the override takes effect depends on the downstream API server's +/// handling of duplicate JSON keys (most Python/Go servers use last-key-wins, +/// but this is not guaranteed by the JSON spec). The `effective_model_name()` +/// trait method should be consulted to determine the model actually used. fn inject_model_override(rig_req: &mut RigRequest, model_override: Option<&str>) { let Some(model) = model_override else { return; @@ -1515,4 +1518,51 @@ mod tests { "different raw IDs should produce different hashed IDs" ); } + + fn make_rig_request(additional_params: Option) -> RigRequest { + RigRequest { + preamble: None, + chat_history: OneOrMany::one(RigMessage::user("test")), + documents: Vec::new(), + tools: Vec::new(), + temperature: None, + max_tokens: None, + tool_choice: None, + additional_params, + } + } + + #[test] + fn test_inject_model_override_creates_params_when_none() { + let mut req = make_rig_request(None); + inject_model_override(&mut req, Some("test-model")); + + let params = req + .additional_params + .expect("additional_params should be Some"); + assert_eq!(params, serde_json::json!({ "model": "test-model" })); + } + + #[test] + fn test_inject_model_override_preserves_existing_params() { + let mut req = make_rig_request(Some(serde_json::json!({ + "cache_control": { "type": "ephemeral" }, + }))); + inject_model_override(&mut req, Some("override-model")); + + let params = req.additional_params.expect("should remain Some"); + let obj = params.as_object().expect("should be object"); + assert_eq!( + obj.get("cache_control"), + Some(&serde_json::json!({ "type": "ephemeral" })) + ); + assert_eq!(obj.get("model"), Some(&serde_json::json!("override-model"))); + } + + #[test] + fn test_inject_model_override_noop_when_none() { + let mut req = make_rig_request(None); + inject_model_override(&mut req, None); + assert!(req.additional_params.is_none()); + } } diff --git a/src/tenant.rs b/src/tenant.rs index 19b0946f..292e5e1a 100644 --- a/src/tenant.rs +++ b/src/tenant.rs @@ -330,35 +330,62 @@ impl TenantScope { /// Add a message to a conversation owned by this tenant. /// - /// Verifies the conversation belongs to this user before adding. + /// Returns `NotFound` if the conversation does not belong to this user. pub async fn add_conversation_message( &self, conversation_id: Uuid, role: &str, content: &str, ) -> Result { + if !self.conversation_belongs_to_user(conversation_id).await? { + return Err(DatabaseError::NotFound { + entity: "conversation".to_string(), + id: conversation_id.to_string(), + }); + } self.inner .add_conversation_message(conversation_id, role, content) .await } + /// Touch a conversation timestamp. Returns `NotFound` if not owned by this user. pub async fn touch_conversation(&self, id: Uuid) -> Result<(), DatabaseError> { + if !self.conversation_belongs_to_user(id).await? { + return Err(DatabaseError::NotFound { + entity: "conversation".to_string(), + id: id.to_string(), + }); + } self.inner.touch_conversation(id).await } + /// List messages in a conversation. Returns `NotFound` if not owned by this user. pub async fn list_conversation_messages( &self, conversation_id: Uuid, ) -> Result, DatabaseError> { + if !self.conversation_belongs_to_user(conversation_id).await? { + return Err(DatabaseError::NotFound { + entity: "conversation".to_string(), + id: conversation_id.to_string(), + }); + } self.inner.list_conversation_messages(conversation_id).await } + /// Paginated message listing. Returns `NotFound` if not owned by this user. pub async fn list_conversation_messages_paginated( &self, conversation_id: Uuid, before: Option>, limit: i64, ) -> Result<(Vec, bool), DatabaseError> { + if !self.conversation_belongs_to_user(conversation_id).await? { + return Err(DatabaseError::NotFound { + entity: "conversation".to_string(), + id: conversation_id.to_string(), + }); + } self.inner .list_conversation_messages_paginated(conversation_id, before, limit) .await @@ -374,21 +401,35 @@ impl TenantScope { .await } + /// Update metadata on a conversation. Returns `NotFound` if not owned by this user. pub async fn update_conversation_metadata_field( &self, id: Uuid, key: &str, value: &serde_json::Value, ) -> Result<(), DatabaseError> { + if !self.conversation_belongs_to_user(id).await? { + return Err(DatabaseError::NotFound { + entity: "conversation".to_string(), + id: id.to_string(), + }); + } self.inner .update_conversation_metadata_field(id, key, value) .await } + /// Get conversation metadata. Returns `NotFound` if not owned by this user. pub async fn get_conversation_metadata( &self, id: Uuid, ) -> Result, DatabaseError> { + if !self.conversation_belongs_to_user(id).await? { + return Err(DatabaseError::NotFound { + entity: "conversation".to_string(), + id: id.to_string(), + }); + } self.inner.get_conversation_metadata(id).await } }