Persist provider name and model config in the session (#5419)

This commit is contained in:
Will Pfleger
2025-11-20 15:55:32 -05:00
committed by GitHub
parent f4724cbf23
commit 682e315be8
18 changed files with 311 additions and 106 deletions
+16 -19
View File
@@ -10,15 +10,21 @@ use goose::session::SessionManager;
use std::path::PathBuf;
#[tokio::main]
async fn main() {
async fn main() -> anyhow::Result<()> {
let _ = dotenv();
let provider = create_with_named_model("databricks", DATABRICKS_DEFAULT_MODEL)
.await
.expect("Couldn't create provider");
let provider = create_with_named_model("databricks", DATABRICKS_DEFAULT_MODEL).await?;
let agent = Agent::new();
let _ = agent.update_provider(provider).await;
let session = SessionManager::create_session(
PathBuf::default(),
"max-turn-test".to_string(),
SessionType::Hidden,
)
.await?;
let _ = agent.update_provider(provider, &session.id).await;
let config = ExtensionConfig::stdio(
"developer",
@@ -27,21 +33,13 @@ async fn main() {
DEFAULT_EXTENSION_TIMEOUT,
)
.with_args(vec!["mcp", "developer"]);
agent.add_extension(config).await.unwrap();
agent.add_extension(config).await?;
println!("Extensions:");
for extension in agent.list_extensions().await {
println!(" {}", extension);
}
let session = SessionManager::create_session(
PathBuf::default(),
"max-turn-test".to_string(),
SessionType::Hidden,
)
.await
.expect("session manager creation failed");
let session_config = SessionConfig {
id: session.id,
schedule_id: None,
@@ -52,13 +50,12 @@ async fn main() {
let user_message = Message::user()
.with_text("can you summarize the readme.md in this dir using just a haiku?");
let mut stream = agent
.reply(user_message, session_config, None)
.await
.unwrap();
let mut stream = agent.reply(user_message, session_config, None).await?;
while let Some(Ok(AgentEvent::Message(message))) = stream.next().await {
println!("{}", serde_json::to_string_pretty(&message).unwrap());
println!("{}", serde_json::to_string_pretty(&message)?);
println!("\n");
}
Ok(())
}
+14 -4
View File
@@ -3,7 +3,7 @@ use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use anyhow::{anyhow, Result};
use anyhow::{anyhow, Context, Result};
use futures::stream::BoxStream;
use futures::{stream, FutureExt, Stream, StreamExt, TryStreamExt};
use uuid::Uuid;
@@ -1259,13 +1259,23 @@ impl Agent {
prompt_manager.add_system_prompt_extra(instruction);
}
pub async fn update_provider(&self, provider: Arc<dyn Provider>) -> Result<()> {
pub async fn update_provider(
&self,
provider: Arc<dyn Provider>,
session_id: &str,
) -> Result<()> {
let mut current_provider = self.provider.lock().await;
*current_provider = Some(provider.clone());
self.update_router_tool_selector(Some(provider), None)
self.update_router_tool_selector(Some(provider.clone()), None)
.await?;
Ok(())
SessionManager::update_session(session_id)
.provider_name(provider.get_name())
.model_config(provider.get_model_config())
.apply()
.await
.context("Failed to persist provider config to session")
}
pub async fn update_router_tool_selector(
+10 -1
View File
@@ -18,6 +18,8 @@ use crate::providers::toolshim::{
use crate::agents::recipe_tools::dynamic_task_tools::should_enabled_subagents;
use crate::session::SessionManager;
#[cfg(test)]
use crate::session::SessionType;
use rmcp::model::Tool;
fn coerce_value(s: &str, schema: &Value) -> Value {
@@ -439,9 +441,16 @@ mod tests {
) -> anyhow::Result<()> {
let agent = crate::agents::Agent::new();
let session = SessionManager::create_session(
std::path::PathBuf::default(),
"test-prepare-tools".to_string(),
SessionType::Hidden,
)
.await?;
let model_config = ModelConfig::new("test-model").unwrap();
let provider = std::sync::Arc::new(MockProvider { model_config });
agent.update_provider(provider).await?;
agent.update_provider(provider, &session.id).await?;
// Disable the router to trigger sorting
agent.disable_router_for_recipe().await;
+1 -1
View File
@@ -117,7 +117,7 @@ fn get_agent_messages(
.map_err(|e| anyhow!("Failed to get sub agent session file path: {}", e))?;
agent
.update_provider(task_config.provider)
.update_provider(task_config.provider, &session_id)
.await
.map_err(|e| anyhow!("Failed to set provider on sub agent: {}", e))?;
+3 -1
View File
@@ -87,7 +87,9 @@ impl AgentManager {
})
.await;
if let Some(provider) = &*self.default_provider.read().await {
agent.update_provider(Arc::clone(provider)).await?;
agent
.update_provider(Arc::clone(provider), &session_id)
.await?;
}
let mut sessions = self.sessions.write().await;
+2 -1
View File
@@ -1,6 +1,7 @@
use once_cell::sync::Lazy;
use serde::{Deserialize, Serialize};
use thiserror::Error;
use utoipa::ToSchema;
const DEFAULT_CONTEXT_LIMIT: usize = 128_000;
@@ -67,7 +68,7 @@ static MODEL_SPECIFIC_LIMITS: Lazy<Vec<(&'static str, usize)>> = Lazy::new(|| {
]
});
#[derive(Debug, Clone, Serialize, Deserialize)]
#[derive(Debug, Clone, Serialize, Deserialize, ToSchema)]
pub struct ModelConfig {
pub model_name: String,
pub context_limit: Option<usize>,
+2 -2
View File
@@ -735,8 +735,6 @@ async fn execute_job(
}
}
agent.update_provider(agent_provider).await?;
let session = SessionManager::create_session(
std::env::current_dir()?,
format!("Scheduled job: {}", job.id),
@@ -744,6 +742,8 @@ async fn execute_job(
)
.await?;
agent.update_provider(agent_provider, &session.id).await?;
let mut jobs_guard = jobs.lock().await;
if let Some((_, job_def)) = jobs_guard.get_mut(job_id.as_str()) {
job_def.current_session_id = Some(session.id.clone());
+70 -5
View File
@@ -1,6 +1,7 @@
use crate::config::paths::Paths;
use crate::conversation::message::Message;
use crate::conversation::Conversation;
use crate::model::ModelConfig;
use crate::providers::base::{Provider, MSG_COUNT_FOR_SESSION_NAME_GENERATION};
use crate::recipe::Recipe;
use crate::session::extension_data::ExtensionData;
@@ -18,7 +19,7 @@ use tokio::sync::OnceCell;
use tracing::{info, warn};
use utoipa::ToSchema;
const CURRENT_SCHEMA_VERSION: i32 = 5;
const CURRENT_SCHEMA_VERSION: i32 = 6;
pub const SESSIONS_FOLDER: &str = "sessions";
pub const DB_NAME: &str = "sessions.db";
@@ -89,6 +90,8 @@ pub struct Session {
pub user_recipe_values: Option<HashMap<String, String>>,
pub conversation: Option<Conversation>,
pub message_count: usize,
pub provider_name: Option<String>,
pub model_config: Option<ModelConfig>,
}
pub struct SessionUpdateBuilder {
@@ -107,6 +110,8 @@ pub struct SessionUpdateBuilder {
schedule_id: Option<Option<String>>,
recipe: Option<Option<Recipe>>,
user_recipe_values: Option<Option<HashMap<String, String>>>,
provider_name: Option<Option<String>>,
model_config: Option<Option<ModelConfig>>,
}
#[derive(Serialize, ToSchema, Debug)]
@@ -134,6 +139,8 @@ impl SessionUpdateBuilder {
schedule_id: None,
recipe: None,
user_recipe_values: None,
provider_name: None,
model_config: None,
}
}
@@ -218,6 +225,16 @@ impl SessionUpdateBuilder {
self
}
pub fn provider_name(mut self, provider_name: impl Into<String>) -> Self {
self.provider_name = Some(Some(provider_name.into()));
self
}
pub fn model_config(mut self, model_config: ModelConfig) -> Self {
self.model_config = Some(Some(model_config));
self
}
pub async fn apply(self) -> Result<()> {
SessionManager::apply_update(self).await
}
@@ -375,6 +392,8 @@ impl Default for Session {
user_recipe_values: None,
conversation: None,
message_count: 0,
provider_name: None,
model_config: None,
}
}
}
@@ -397,6 +416,9 @@ impl sqlx::FromRow<'_, sqlx::sqlite::SqliteRow> for Session {
let user_recipe_values =
user_recipe_values_json.and_then(|json| serde_json::from_str(&json).ok());
let model_config_json: Option<String> = row.try_get("model_config_json").ok().flatten();
let model_config = model_config_json.and_then(|json| serde_json::from_str(&json).ok());
let name: String = {
let name_val: String = row.try_get("name").unwrap_or_default();
if !name_val.is_empty() {
@@ -434,6 +456,8 @@ impl sqlx::FromRow<'_, sqlx::sqlite::SqliteRow> for Session {
user_recipe_values,
conversation: None,
message_count: row.try_get("message_count").unwrap_or(0) as usize,
provider_name: row.try_get("provider_name").ok().flatten(),
model_config,
})
}
}
@@ -521,7 +545,9 @@ impl SessionStorage {
accumulated_output_tokens INTEGER,
schedule_id TEXT,
recipe_json TEXT,
user_recipe_values_json TEXT
user_recipe_values_json TEXT,
provider_name TEXT,
model_config_json TEXT
)
"#,
)
@@ -618,14 +644,20 @@ impl SessionStorage {
None => None,
};
let model_config_json = match &session.model_config {
Some(model_config) => Some(serde_json::to_string(model_config)?),
None => None,
};
sqlx::query(
r#"
INSERT INTO sessions (
id, name, user_set_name, session_type, working_dir, created_at, updated_at, extension_data,
total_tokens, input_tokens, output_tokens,
accumulated_total_tokens, accumulated_input_tokens, accumulated_output_tokens,
schedule_id, recipe_json, user_recipe_values_json
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
schedule_id, recipe_json, user_recipe_values_json,
provider_name, model_config_json
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&session.id)
@@ -645,6 +677,8 @@ impl SessionStorage {
.bind(&session.schedule_id)
.bind(recipe_json)
.bind(user_recipe_values_json)
.bind(&session.provider_name)
.bind(model_config_json)
.execute(&mut *tx)
.await?;
@@ -771,6 +805,23 @@ impl SessionStorage {
.execute(&self.pool)
.await?;
}
6 => {
sqlx::query(
r#"
ALTER TABLE sessions ADD COLUMN provider_name TEXT
"#,
)
.execute(&self.pool)
.await?;
sqlx::query(
r#"
ALTER TABLE sessions ADD COLUMN model_config_json TEXT
"#,
)
.execute(&self.pool)
.await?;
}
_ => {
anyhow::bail!("Unknown migration version: {}", version);
}
@@ -824,7 +875,8 @@ impl SessionStorage {
SELECT id, working_dir, name, description, user_set_name, session_type, created_at, updated_at, extension_data,
total_tokens, input_tokens, output_tokens,
accumulated_total_tokens, accumulated_input_tokens, accumulated_output_tokens,
schedule_id, recipe_json, user_recipe_values_json
schedule_id, recipe_json, user_recipe_values_json,
provider_name, model_config_json
FROM sessions
WHERE id = ?
"#,
@@ -850,6 +902,7 @@ impl SessionStorage {
Ok(session)
}
#[allow(clippy::too_many_lines)]
async fn apply_update(&self, builder: SessionUpdateBuilder) -> Result<()> {
let mut updates = Vec::new();
let mut query = String::from("UPDATE sessions SET ");
@@ -884,6 +937,8 @@ impl SessionStorage {
add_update!(builder.schedule_id, "schedule_id");
add_update!(builder.recipe, "recipe_json");
add_update!(builder.user_recipe_values, "user_recipe_values_json");
add_update!(builder.provider_name, "provider_name");
add_update!(builder.model_config, "model_config_json");
if updates.is_empty() {
return Ok(());
@@ -940,6 +995,15 @@ impl SessionStorage {
.transpose()?;
q = q.bind(user_recipe_values_json);
}
if let Some(provider_name) = builder.provider_name {
q = q.bind(provider_name);
}
if let Some(model_config) = builder.model_config {
let model_config_json = model_config
.map(|mc| serde_json::to_string(&mc))
.transpose()?;
q = q.bind(model_config_json);
}
let mut tx = self.pool.begin().await?;
q = q.bind(&builder.session_id);
@@ -1050,6 +1114,7 @@ impl SessionStorage {
s.total_tokens, s.input_tokens, s.output_tokens,
s.accumulated_total_tokens, s.accumulated_input_tokens, s.accumulated_output_tokens,
s.schedule_id, s.recipe_json, s.user_recipe_values_json,
s.provider_name, s.model_config_json,
COUNT(m.id) as message_count
FROM sessions s
INNER JOIN messages m ON s.id = m.session_id
+3 -1
View File
@@ -373,7 +373,6 @@ mod tests {
async fn test_max_turns_limit() -> Result<()> {
let agent = Agent::new();
let provider = Arc::new(MockToolProvider::new());
agent.update_provider(provider).await?;
let user_message = Message::user().with_text("Hello");
let session = SessionManager::create_session(
@@ -382,6 +381,9 @@ mod tests {
SessionType::Hidden,
)
.await?;
agent.update_provider(provider, &session.id).await?;
let session_config = SessionConfig {
id: session.id,
schedule_id: None,