Move token limits to backend (#2484)
This commit is contained in:
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user