Move token limits to backend (#2484)

This commit is contained in:
Zane
2025-05-12 13:52:46 -07:00
committed by GitHub
parent b4aadeb5d4
commit 7027de6238
14 changed files with 260 additions and 1121 deletions
@@ -10,6 +10,7 @@ use etcetera::{choose_app_strategy, AppStrategy, AppStrategyArgs};
use goose::config::Config;
use goose::config::{extensions::name_to_key, PermissionManager};
use goose::config::{ExtensionConfigManager, ExtensionEntry};
use goose::model::ModelConfig;
use goose::providers::base::ProviderMetadata;
use goose::providers::providers as get_providers;
use goose::{agents::ExtensionConfig, config::permission::PermissionLevel};
@@ -154,6 +155,14 @@ pub async fn read_config(
) -> Result<Json<Value>, StatusCode> {
verify_secret_key(&headers, &state)?;
// Special handling for model-limits
if query.key == "model-limits" {
let limits = ModelConfig::get_all_model_limits();
return Ok(Json(
serde_json::to_value(limits).map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?,
));
}
let config = Config::global();
match config.get(&query.key, query.is_secret) {
@@ -481,3 +490,45 @@ pub fn routes(state: Arc<AppState>) -> Router {
.route("/config/permissions", post(upsert_permissions))
.with_state(state)
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_read_model_limits() {
// Create test state and headers
let test_state = AppState::new(
Arc::new(goose::agents::Agent::default()),
"test".to_string(),
)
.await;
let mut headers = HeaderMap::new();
headers.insert("X-Secret-Key", "test".parse().unwrap());
// Execute
let result = read_config(
State(test_state),
headers,
Json(ConfigKeyQuery {
key: "model-limits".to_string(),
is_secret: false,
}),
)
.await;
// Assert
assert!(result.is_ok());
let response = result.unwrap();
// Parse the response and check the contents
let limits: Vec<goose::model::ModelLimitConfig> =
serde_json::from_value(response.0).unwrap();
assert!(!limits.is_empty());
// Check for some expected patterns
let gpt4_limit = limits.iter().find(|l| l.pattern == "gpt-4o");
assert!(gpt4_limit.is_some());
assert_eq!(gpt4_limit.unwrap().context_limit, 128_000);
}
}
-229
View File
@@ -1,229 +0,0 @@
use super::utils::verify_secret_key;
use crate::state::AppState;
use axum::{
extract::{Query, State},
routing::{delete, get, post},
Json, Router,
};
use goose::config::Config;
use http::{HeaderMap, StatusCode};
use once_cell::sync::Lazy;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::{collections::HashMap, sync::Arc};
#[derive(Serialize)]
struct ConfigResponse {
error: bool,
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct ConfigRequest {
key: String,
value: String,
is_secret: bool,
}
async fn store_config(
State(state): State<Arc<AppState>>,
headers: HeaderMap,
Json(request): Json<ConfigRequest>,
) -> Result<Json<ConfigResponse>, StatusCode> {
verify_secret_key(&headers, &state)?;
let config = Config::global();
let result = if request.is_secret {
config.set_secret(&request.key, Value::String(request.value))
} else {
config.set_param(&request.key, Value::String(request.value))
};
match result {
Ok(_) => Ok(Json(ConfigResponse { error: false })),
Err(_) => Ok(Json(ConfigResponse { error: true })),
}
}
#[derive(Debug, Serialize, Deserialize)]
pub struct ProviderConfigRequest {
pub providers: Vec<String>,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct ConfigStatus {
pub is_set: bool,
pub location: Option<String>,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct ProviderResponse {
pub supported: bool,
pub name: Option<String>,
pub description: Option<String>,
pub models: Option<Vec<String>>,
pub config_status: HashMap<String, ConfigStatus>,
}
#[derive(Debug, Serialize, Deserialize)]
struct ProviderConfig {
name: String,
description: String,
models: Vec<String>,
required_keys: Vec<String>,
}
static PROVIDER_ENV_REQUIREMENTS: Lazy<HashMap<String, ProviderConfig>> = Lazy::new(|| {
let contents = include_str!("providers_and_keys.json");
serde_json::from_str(contents).expect("Failed to parse providers_and_keys.json")
});
fn check_key_status(config: &Config, key: &str) -> (bool, Option<String>) {
if let Ok(_value) = std::env::var(key) {
(true, Some("env".to_string()))
} else if config.get_param::<String>(key).is_ok() {
(true, Some("yaml".to_string()))
} else if config.get_secret::<String>(key).is_ok() {
(true, Some("keyring".to_string()))
} else {
(false, None)
}
}
async fn check_provider_configs(
Json(request): Json<ProviderConfigRequest>,
) -> Result<Json<HashMap<String, ProviderResponse>>, StatusCode> {
let mut response = HashMap::new();
let config = Config::global();
for provider_name in request.providers {
if let Some(provider_config) = PROVIDER_ENV_REQUIREMENTS.get(&provider_name) {
let mut config_status = HashMap::new();
for key in &provider_config.required_keys {
let (key_set, key_location) = check_key_status(config, key);
config_status.insert(
key.to_string(),
ConfigStatus {
is_set: key_set,
location: key_location,
},
);
}
response.insert(
provider_name,
ProviderResponse {
supported: true,
name: Some(provider_config.name.clone()),
description: Some(provider_config.description.clone()),
models: Some(provider_config.models.clone()),
config_status,
},
);
} else {
response.insert(
provider_name,
ProviderResponse {
supported: false,
name: None,
description: None,
models: None,
config_status: HashMap::new(),
},
);
}
}
Ok(Json(response))
}
#[derive(Deserialize)]
pub struct GetConfigQuery {
key: String,
}
#[derive(Serialize)]
pub struct GetConfigResponse {
value: Option<String>,
}
pub async fn get_config(
State(state): State<Arc<AppState>>,
headers: HeaderMap,
Query(query): Query<GetConfigQuery>,
) -> Result<Json<GetConfigResponse>, StatusCode> {
verify_secret_key(&headers, &state)?;
// Fetch the configuration value. Right now we don't allow get a secret.
let config = Config::global();
let value = if let Ok(config_value) = config.get_param::<String>(&query.key) {
Some(config_value)
} else {
std::env::var(&query.key).ok()
};
// Return the value
Ok(Json(GetConfigResponse { value }))
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct DeleteConfigRequest {
key: String,
is_secret: bool,
}
async fn delete_config(
State(state): State<Arc<AppState>>,
headers: HeaderMap,
Json(request): Json<DeleteConfigRequest>,
) -> Result<StatusCode, StatusCode> {
verify_secret_key(&headers, &state)?;
// Attempt to delete the key
let config = Config::global();
let result = if request.is_secret {
config.delete_secret(&request.key)
} else {
config.delete(&request.key)
};
match result {
Ok(_) => Ok(StatusCode::NO_CONTENT),
Err(_) => Err(StatusCode::NOT_FOUND),
}
}
pub fn routes(state: Arc<AppState>) -> Router {
Router::new()
.route("/configs/providers", post(check_provider_configs))
.route("/configs/get", get(get_config))
.route("/configs/store", post(store_config))
.route("/configs/delete", delete(delete_config))
.with_state(state)
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_unsupported_provider() {
// Setup
let request = ProviderConfigRequest {
providers: vec!["unsupported_provider".to_string()],
};
// Execute
let result = check_provider_configs(Json(request)).await;
// Assert
assert!(result.is_ok());
let Json(response) = result.unwrap();
let provider_status = response
.get("unsupported_provider")
.expect("Provider should exist");
assert!(!provider_status.supported);
assert!(provider_status.config_status.is_empty());
}
}
-2
View File
@@ -1,7 +1,6 @@
// Export route modules
pub mod agent;
pub mod config_management;
pub mod configs;
pub mod context;
pub mod extension;
pub mod health;
@@ -21,7 +20,6 @@ pub fn configure(state: Arc<crate::state::AppState>) -> Router {
.merge(agent::routes(state.clone()))
.merge(context::routes(state.clone()))
.merge(extension::routes(state.clone()))
.merge(configs::routes(state.clone()))
.merge(config_management::routes(state.clone()))
.merge(recipe::routes(state.clone()))
.merge(session::routes(state.clone()))
+62 -22
View File
@@ -1,4 +1,6 @@
use once_cell::sync::Lazy;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
const DEFAULT_CONTEXT_LIMIT: usize = 128_000;
@@ -6,6 +8,32 @@ const DEFAULT_CONTEXT_LIMIT: usize = 128_000;
pub const GPT_4O_TOKENIZER: &str = "Xenova--gpt-4o";
pub const CLAUDE_TOKENIZER: &str = "Xenova--claude-tokenizer";
// Define the model limits as a static HashMap for reuse
static MODEL_SPECIFIC_LIMITS: Lazy<HashMap<&'static str, usize>> = Lazy::new(|| {
let mut map = HashMap::new();
// OpenAI models, https://platform.openai.com/docs/models#models-overview
map.insert("gpt-4o", 128_000);
map.insert("gpt-4-turbo", 128_000);
map.insert("o1-mini", 128_000);
map.insert("o1-preview", 128_000);
map.insert("o1", 200_000);
map.insert("o3-mini", 200_000);
map.insert("gpt-4.1", 1_000_000);
map.insert("gpt-4-1", 1_000_000);
// Anthropic models, https://docs.anthropic.com/en/docs/about-claude/models
map.insert("claude-3", 200_000);
// Google models, https://ai.google/get-started/our-models/
map.insert("gemini-2.5", 1_000_000);
map.insert("gemini-2-5", 1_000_000);
// Meta Llama models, https://github.com/meta-llama/llama-models/tree/main?tab=readme-ov-file#llama-models-1
map.insert("llama3.2", 128_000);
map.insert("llama3.3", 128_000);
map
});
/// Configuration for model-specific settings and limits
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelConfig {
@@ -27,6 +55,13 @@ pub struct ModelConfig {
pub toolshim_model: Option<String>,
}
/// Struct to represent model pattern matches and their limits
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelLimitConfig {
pub pattern: String,
pub context_limit: usize,
}
impl ModelConfig {
/// Create a new ModelConfig with the specified model name
///
@@ -70,29 +105,23 @@ impl ModelConfig {
/// Get model-specific context limit based on model name
fn get_model_specific_limit(model_name: &str) -> Option<usize> {
// Implement some sensible defaults
match model_name {
// OpenAI models, https://platform.openai.com/docs/models#models-overview
name if name.contains("gpt-4o") => Some(128_000),
name if name.contains("gpt-4-turbo") => Some(128_000),
name if name.contains("o1-mini") || name.contains("o1-preview") => Some(128_000),
name if name.contains("o1") => Some(200_000),
name if name.contains("o3-mini") => Some(200_000),
name if name.contains("gpt-4.1") => Some(1_000_000),
name if name.contains("gpt-4-1") => Some(1_000_000),
// Anthropic models, https://docs.anthropic.com/en/docs/about-claude/models
name if name.contains("claude-3") => Some(200_000),
// Google models, https://ai.google/get-started/our-models/
name if name.contains("gemini-2.5") => Some(1_000_000),
name if name.contains("gemini-2-5") => Some(1_000_000),
// Meta Llama models, https://github.com/meta-llama/llama-models/tree/main?tab=readme-ov-file#llama-models-1
name if name.contains("llama3.2") => Some(128_000),
name if name.contains("llama3.3") => Some(128_000),
_ => None,
for (pattern, &limit) in MODEL_SPECIFIC_LIMITS.iter() {
if model_name.contains(pattern) {
return Some(limit);
}
}
None
}
/// Get all model pattern matches and their limits
pub fn get_all_model_limits() -> Vec<ModelLimitConfig> {
MODEL_SPECIFIC_LIMITS
.iter()
.map(|(&pattern, &context_limit)| ModelLimitConfig {
pattern: pattern.to_string(),
context_limit,
})
.collect()
}
/// Set an explicit context limit
@@ -215,4 +244,15 @@ mod tests {
let config = ModelConfig::new("test-model".to_string());
assert_eq!(config.temperature, None);
}
#[test]
fn test_get_all_model_limits() {
let limits = ModelConfig::get_all_model_limits();
assert!(!limits.is_empty());
// Test that we can find specific patterns
let gpt4_limit = limits.iter().find(|l| l.pattern == "gpt-4o");
assert!(gpt4_limit.is_some());
assert_eq!(gpt4_limit.unwrap().context_limit, 128_000);
}
}