Chat bottom menu bar token and tools alerts (#2146)
Co-authored-by: Lily Delalande <119957291+lily-de@users.noreply.github.com>
This commit is contained in:
@@ -4,8 +4,7 @@ use goose::agents::ExtensionConfig;
|
||||
use goose::config::permission::PermissionLevel;
|
||||
use goose::config::ExtensionEntry;
|
||||
use goose::permission::permission_confirmation::PrincipalType;
|
||||
use goose::providers::base::ConfigKey;
|
||||
use goose::providers::base::ProviderMetadata;
|
||||
use goose::providers::base::{ConfigKey, ModelInfo, ProviderMetadata};
|
||||
use mcp_core::tool::{Tool, ToolAnnotations};
|
||||
use utoipa::OpenApi;
|
||||
|
||||
@@ -47,6 +46,7 @@ use utoipa::OpenApi;
|
||||
ToolInfo,
|
||||
PermissionLevel,
|
||||
PrincipalType,
|
||||
ModelInfo,
|
||||
))
|
||||
)]
|
||||
pub struct ApiDoc;
|
||||
|
||||
@@ -124,10 +124,7 @@ impl Provider for AnthropicProvider {
|
||||
"Anthropic",
|
||||
"Claude and other models from Anthropic",
|
||||
ANTHROPIC_DEFAULT_MODEL,
|
||||
ANTHROPIC_KNOWN_MODELS
|
||||
.iter()
|
||||
.map(|&s| s.to_string())
|
||||
.collect(),
|
||||
ANTHROPIC_KNOWN_MODELS.to_vec(),
|
||||
ANTHROPIC_DOC_URL,
|
||||
vec![
|
||||
ConfigKey::new("ANTHROPIC_API_KEY", true, true, None),
|
||||
|
||||
@@ -101,10 +101,7 @@ impl Provider for AzureProvider {
|
||||
"Azure OpenAI",
|
||||
"Models through Azure OpenAI Service",
|
||||
"gpt-4o",
|
||||
AZURE_OPENAI_KNOWN_MODELS
|
||||
.iter()
|
||||
.map(|s| s.to_string())
|
||||
.collect(),
|
||||
AZURE_OPENAI_KNOWN_MODELS.to_vec(),
|
||||
AZURE_DOC_URL,
|
||||
vec![
|
||||
ConfigKey::new("AZURE_OPENAI_API_KEY", true, true, None),
|
||||
|
||||
@@ -25,6 +25,15 @@ pub fn get_current_model() -> Option<String> {
|
||||
CURRENT_MODEL.lock().ok().and_then(|model| model.clone())
|
||||
}
|
||||
|
||||
/// Information about a model's capabilities
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, ToSchema, PartialEq)]
|
||||
pub struct ModelInfo {
|
||||
/// The name of the model
|
||||
pub name: String,
|
||||
/// The maximum context length this model supports
|
||||
pub context_limit: usize,
|
||||
}
|
||||
|
||||
/// Metadata about a provider's configuration requirements and capabilities
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, ToSchema)]
|
||||
pub struct ProviderMetadata {
|
||||
@@ -36,9 +45,9 @@ pub struct ProviderMetadata {
|
||||
pub description: String,
|
||||
/// The default/recommended model for this provider
|
||||
pub default_model: String,
|
||||
/// A list of currently known models
|
||||
/// A list of currently known models with their capabilities
|
||||
/// TODO: eventually query the apis directly
|
||||
pub known_models: Vec<String>,
|
||||
pub known_models: Vec<ModelInfo>,
|
||||
/// Link to the docs where models can be found
|
||||
pub model_doc_link: String,
|
||||
/// Required configuration keys
|
||||
@@ -51,7 +60,7 @@ impl ProviderMetadata {
|
||||
display_name: &str,
|
||||
description: &str,
|
||||
default_model: &str,
|
||||
known_models: Vec<String>,
|
||||
model_names: Vec<&str>,
|
||||
model_doc_link: &str,
|
||||
config_keys: Vec<ConfigKey>,
|
||||
) -> Self {
|
||||
@@ -60,7 +69,13 @@ impl ProviderMetadata {
|
||||
display_name: display_name.to_string(),
|
||||
description: description.to_string(),
|
||||
default_model: default_model.to_string(),
|
||||
known_models,
|
||||
known_models: model_names
|
||||
.iter()
|
||||
.map(|&name| ModelInfo {
|
||||
name: name.to_string(),
|
||||
context_limit: ModelConfig::new(name.to_string()).context_limit(),
|
||||
})
|
||||
.collect(),
|
||||
model_doc_link: model_doc_link.to_string(),
|
||||
config_keys,
|
||||
}
|
||||
@@ -168,6 +183,7 @@ pub trait Provider: Send + Sync {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::collections::HashMap;
|
||||
|
||||
use serde_json::json;
|
||||
|
||||
@@ -214,4 +230,61 @@ mod tests {
|
||||
let model = get_current_model();
|
||||
assert_eq!(model, Some("claude-3.5-sonnet".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_provider_metadata_context_limits() {
|
||||
// Test that ProviderMetadata::new correctly sets context limits
|
||||
let test_models = vec!["gpt-4o", "claude-3-5-sonnet-latest", "unknown-model"];
|
||||
let metadata = ProviderMetadata::new(
|
||||
"test",
|
||||
"Test Provider",
|
||||
"Test Description",
|
||||
"gpt-4o",
|
||||
test_models,
|
||||
"https://example.com",
|
||||
vec![],
|
||||
);
|
||||
|
||||
let model_info: HashMap<String, usize> = metadata
|
||||
.known_models
|
||||
.into_iter()
|
||||
.map(|m| (m.name, m.context_limit))
|
||||
.collect();
|
||||
|
||||
// gpt-4o should have 128k limit
|
||||
assert_eq!(*model_info.get("gpt-4o").unwrap(), 128_000);
|
||||
|
||||
// claude-3-5-sonnet-latest should have 200k limit
|
||||
assert_eq!(
|
||||
*model_info.get("claude-3-5-sonnet-latest").unwrap(),
|
||||
200_000
|
||||
);
|
||||
|
||||
// unknown model should have default limit (128k)
|
||||
assert_eq!(*model_info.get("unknown-model").unwrap(), 128_000);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_model_info_creation() {
|
||||
// Test direct ModelInfo creation
|
||||
let info = ModelInfo {
|
||||
name: "test-model".to_string(),
|
||||
context_limit: 1000,
|
||||
};
|
||||
assert_eq!(info.context_limit, 1000);
|
||||
|
||||
// Test equality
|
||||
let info2 = ModelInfo {
|
||||
name: "test-model".to_string(),
|
||||
context_limit: 1000,
|
||||
};
|
||||
assert_eq!(info, info2);
|
||||
|
||||
// Test inequality
|
||||
let info3 = ModelInfo {
|
||||
name: "test-model".to_string(),
|
||||
context_limit: 2000,
|
||||
};
|
||||
assert_ne!(info, info3);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -85,7 +85,7 @@ impl Provider for BedrockProvider {
|
||||
"Amazon Bedrock",
|
||||
"Run models through Amazon Bedrock. You may have to set 'AWS_' environment variables to configure authentication.",
|
||||
BEDROCK_DEFAULT_MODEL,
|
||||
BEDROCK_KNOWN_MODELS.iter().map(|s| s.to_string()).collect(),
|
||||
BEDROCK_KNOWN_MODELS.to_vec(),
|
||||
BEDROCK_DOC_LINK,
|
||||
vec![ConfigKey::new("AWS_PROFILE", true, false, Some("default"))],
|
||||
)
|
||||
|
||||
@@ -248,10 +248,7 @@ impl Provider for DatabricksProvider {
|
||||
"Databricks",
|
||||
"Models on Databricks AI Gateway",
|
||||
DATABRICKS_DEFAULT_MODEL,
|
||||
DATABRICKS_KNOWN_MODELS
|
||||
.iter()
|
||||
.map(|&s| s.to_string())
|
||||
.collect(),
|
||||
DATABRICKS_KNOWN_MODELS.to_vec(),
|
||||
DATABRICKS_DOC_URL,
|
||||
vec![
|
||||
ConfigKey::new("DATABRICKS_HOST", true, false, None),
|
||||
|
||||
@@ -425,7 +425,7 @@ impl Provider for GcpVertexAIProvider {
|
||||
where
|
||||
Self: Sized,
|
||||
{
|
||||
let known_models = vec![
|
||||
let model_strings: Vec<String> = vec![
|
||||
GcpVertexAIModel::Claude(ClaudeVersion::Sonnet35),
|
||||
GcpVertexAIModel::Claude(ClaudeVersion::Sonnet35V2),
|
||||
GcpVertexAIModel::Claude(ClaudeVersion::Sonnet37),
|
||||
@@ -434,10 +434,12 @@ impl Provider for GcpVertexAIProvider {
|
||||
GcpVertexAIModel::Gemini(GeminiVersion::Flash20),
|
||||
GcpVertexAIModel::Gemini(GeminiVersion::Pro20Exp),
|
||||
]
|
||||
.into_iter()
|
||||
.iter()
|
||||
.map(|model| model.to_string())
|
||||
.collect();
|
||||
|
||||
let known_models: Vec<&str> = model_strings.iter().map(|s| s.as_str()).collect();
|
||||
|
||||
ProviderMetadata::new(
|
||||
"gcp_vertex_ai",
|
||||
"GCP Vertex AI",
|
||||
@@ -583,12 +585,13 @@ mod tests {
|
||||
#[test]
|
||||
fn test_provider_metadata() {
|
||||
let metadata = GcpVertexAIProvider::metadata();
|
||||
assert!(metadata
|
||||
let model_names: Vec<String> = metadata
|
||||
.known_models
|
||||
.contains(&"claude-3-5-sonnet-v2@20241022".to_string()));
|
||||
assert!(metadata
|
||||
.known_models
|
||||
.contains(&"gemini-1.5-pro-002".to_string()));
|
||||
.iter()
|
||||
.map(|m| m.name.clone())
|
||||
.collect();
|
||||
assert!(model_names.contains(&"claude-3-5-sonnet-v2@20241022".to_string()));
|
||||
assert!(model_names.contains(&"gemini-1.5-pro-002".to_string()));
|
||||
// Should contain the original 2 config keys plus 4 new retry-related ones
|
||||
assert_eq!(metadata.config_keys.len(), 6);
|
||||
}
|
||||
|
||||
@@ -129,7 +129,7 @@ impl Provider for GoogleProvider {
|
||||
"Google Gemini",
|
||||
"Gemini models from Google AI",
|
||||
GOOGLE_DEFAULT_MODEL,
|
||||
GOOGLE_KNOWN_MODELS.iter().map(|&s| s.to_string()).collect(),
|
||||
GOOGLE_KNOWN_MODELS.to_vec(),
|
||||
GOOGLE_DOC_URL,
|
||||
vec![
|
||||
ConfigKey::new("GOOGLE_API_KEY", true, true, None),
|
||||
|
||||
@@ -105,7 +105,7 @@ impl Provider for GroqProvider {
|
||||
"Groq",
|
||||
"Fast inference with Groq hardware",
|
||||
GROQ_DEFAULT_MODEL,
|
||||
GROQ_KNOWN_MODELS.iter().map(|&s| s.to_string()).collect(),
|
||||
GROQ_KNOWN_MODELS.to_vec(),
|
||||
GROQ_DOC_URL,
|
||||
vec![
|
||||
ConfigKey::new("GROQ_API_KEY", true, true, None),
|
||||
|
||||
@@ -102,7 +102,7 @@ impl Provider for OllamaProvider {
|
||||
"Ollama",
|
||||
"Local open source models",
|
||||
OLLAMA_DEFAULT_MODEL,
|
||||
OLLAMA_KNOWN_MODELS.iter().map(|&s| s.to_string()).collect(),
|
||||
OLLAMA_KNOWN_MODELS.to_vec(),
|
||||
OLLAMA_DOC_URL,
|
||||
vec![ConfigKey::new(
|
||||
"OLLAMA_HOST",
|
||||
|
||||
@@ -122,10 +122,7 @@ impl Provider for OpenAiProvider {
|
||||
"OpenAI",
|
||||
"GPT-4 and other OpenAI models, including OpenAI compatible ones",
|
||||
OPEN_AI_DEFAULT_MODEL,
|
||||
OPEN_AI_KNOWN_MODELS
|
||||
.iter()
|
||||
.map(|&s| s.to_string())
|
||||
.collect(),
|
||||
OPEN_AI_KNOWN_MODELS.to_vec(),
|
||||
OPEN_AI_DOC_URL,
|
||||
vec![
|
||||
ConfigKey::new("OPENAI_API_KEY", true, true, None),
|
||||
|
||||
@@ -229,10 +229,7 @@ impl Provider for OpenRouterProvider {
|
||||
"OpenRouter",
|
||||
"Router for many model providers",
|
||||
OPENROUTER_DEFAULT_MODEL,
|
||||
OPENROUTER_KNOWN_MODELS
|
||||
.iter()
|
||||
.map(|&s| s.to_string())
|
||||
.collect(),
|
||||
OPENROUTER_KNOWN_MODELS.to_vec(),
|
||||
OPENROUTER_DOC_URL,
|
||||
vec![
|
||||
ConfigKey::new("OPENROUTER_API_KEY", true, true, None),
|
||||
|
||||
Reference in New Issue
Block a user