Add PKCE support for Tetrate Agent Router Service (#4165)
Signed-off-by: John Landa <jonathanlanda@gmail.com> Co-authored-by: Michael Neale <michael.neale@gmail.com>
This commit is contained in:
@@ -20,6 +20,7 @@ use super::{
|
||||
provider_registry::ProviderRegistry,
|
||||
sagemaker_tgi::SageMakerTgiProvider,
|
||||
snowflake::SnowflakeProvider,
|
||||
tetrate::TetrateProvider,
|
||||
venice::VeniceProvider,
|
||||
xai::XaiProvider,
|
||||
};
|
||||
@@ -55,6 +56,7 @@ static REGISTRY: Lazy<RwLock<ProviderRegistry>> = Lazy::new(|| {
|
||||
registry.register::<OpenRouterProvider, _>(OpenRouterProvider::from_env);
|
||||
registry.register::<SageMakerTgiProvider, _>(SageMakerTgiProvider::from_env);
|
||||
registry.register::<SnowflakeProvider, _>(SnowflakeProvider::from_env);
|
||||
registry.register::<TetrateProvider, _>(TetrateProvider::from_env);
|
||||
registry.register::<VeniceProvider, _>(VeniceProvider::from_env);
|
||||
registry.register::<XaiProvider, _>(XaiProvider::from_env);
|
||||
|
||||
|
||||
@@ -29,6 +29,7 @@ mod retry;
|
||||
pub mod sagemaker_tgi;
|
||||
pub mod snowflake;
|
||||
pub mod testprovider;
|
||||
pub mod tetrate;
|
||||
pub mod toolshim;
|
||||
pub mod usage_estimator;
|
||||
pub mod utils;
|
||||
|
||||
@@ -0,0 +1,264 @@
|
||||
use anyhow::{Error, Result};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::api_client::{ApiClient, AuthMethod};
|
||||
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
use super::errors::ProviderError;
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::{
|
||||
emit_debug_trace, get_model, handle_response_google_compat, handle_response_openai_compat,
|
||||
is_google_model,
|
||||
};
|
||||
use crate::config::signup_tetrate::TETRATE_DEFAULT_MODEL;
|
||||
use crate::conversation::message::Message;
|
||||
use crate::impl_provider_default;
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::formats::openai::{create_request, get_usage, response_to_message};
|
||||
use rmcp::model::Tool;
|
||||
|
||||
// Tetrate Agent Router Service can run many models, we suggest the default
|
||||
pub const TETRATE_KNOWN_MODELS: &[&str] = &[
|
||||
"claude-opus-4-1",
|
||||
"claude-3-7-sonnet-latest",
|
||||
"claude-3-5-sonnet-latest",
|
||||
"claude-3-5-haiku-latest",
|
||||
"gemini-2.5-pro",
|
||||
"gemini-2.0-flash",
|
||||
"gemini-2.0-flash-lite",
|
||||
"gpt-5",
|
||||
"gpt-5-mini",
|
||||
"gpt-5-nano",
|
||||
"gpt-4.1",
|
||||
];
|
||||
pub const TETRATE_DOC_URL: &str = "https://router.tetrate.ai";
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct TetrateProvider {
|
||||
#[serde(skip)]
|
||||
api_client: ApiClient,
|
||||
model: ModelConfig,
|
||||
}
|
||||
|
||||
impl_provider_default!(TetrateProvider);
|
||||
|
||||
impl TetrateProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let api_key: String = config.get_secret("TETRATE_API_KEY")?;
|
||||
// API host for LLM endpoints (/v1/chat/completions, /v1/models)
|
||||
let host: String = config
|
||||
.get_param("TETRATE_HOST")
|
||||
.unwrap_or_else(|_| "https://api.router.tetrate.ai".to_string());
|
||||
|
||||
let auth = AuthMethod::BearerToken(api_key);
|
||||
let api_client = ApiClient::new(host, auth)?
|
||||
.with_header("HTTP-Referer", "https://block.github.io/goose")?
|
||||
.with_header("X-Title", "Goose")?;
|
||||
|
||||
Ok(Self { api_client, model })
|
||||
}
|
||||
|
||||
async fn post(&self, payload: &Value) -> Result<Value, ProviderError> {
|
||||
let response = self
|
||||
.api_client
|
||||
.response_post("v1/chat/completions", payload)
|
||||
.await?;
|
||||
|
||||
// Handle Google-compatible model responses differently
|
||||
if is_google_model(payload) {
|
||||
return handle_response_google_compat(response).await;
|
||||
}
|
||||
|
||||
// For OpenAI-compatible models, parse the response body to JSON
|
||||
let response_body = handle_response_openai_compat(response)
|
||||
.await
|
||||
.map_err(|e| ProviderError::RequestFailed(format!("Failed to parse response: {e}")))?;
|
||||
|
||||
let _debug = format!(
|
||||
"Tetrate Agent Router Service request with payload: {} and response: {}",
|
||||
serde_json::to_string_pretty(payload).unwrap_or_else(|_| "Invalid JSON".to_string()),
|
||||
serde_json::to_string_pretty(&response_body)
|
||||
.unwrap_or_else(|_| "Invalid JSON".to_string())
|
||||
);
|
||||
|
||||
// Tetrate Agent Router Service can return errors in 200 OK responses, so we have to check for errors explicitly
|
||||
if let Some(error_obj) = response_body.get("error") {
|
||||
// If there's an error object, extract the error message and code
|
||||
let error_message = error_obj
|
||||
.get("message")
|
||||
.and_then(|m| m.as_str())
|
||||
.unwrap_or("Unknown Tetrate Agent Router Service error");
|
||||
|
||||
let error_code = error_obj.get("code").and_then(|c| c.as_u64()).unwrap_or(0);
|
||||
|
||||
// Check for context length errors in the error message
|
||||
if error_code == 400 && error_message.contains("maximum context length") {
|
||||
return Err(ProviderError::ContextLengthExceeded(
|
||||
error_message.to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
// Return appropriate error based on the error code
|
||||
match error_code {
|
||||
401 | 403 => return Err(ProviderError::Authentication(error_message.to_string())),
|
||||
429 => return Err(ProviderError::RateLimitExceeded(error_message.to_string())),
|
||||
500 | 503 => return Err(ProviderError::ServerError(error_message.to_string())),
|
||||
_ => return Err(ProviderError::RequestFailed(error_message.to_string())),
|
||||
}
|
||||
}
|
||||
|
||||
// No error detected, return the response body
|
||||
Ok(response_body)
|
||||
}
|
||||
}
|
||||
|
||||
fn create_request_based_on_model(
|
||||
provider: &TetrateProvider,
|
||||
system: &str,
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
) -> anyhow::Result<Value, Error> {
|
||||
let payload = create_request(
|
||||
&provider.model,
|
||||
system,
|
||||
messages,
|
||||
tools,
|
||||
&super::utils::ImageFormat::OpenAi,
|
||||
)?;
|
||||
|
||||
Ok(payload)
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Provider for TetrateProvider {
|
||||
fn metadata() -> ProviderMetadata {
|
||||
ProviderMetadata::new(
|
||||
"tetrate",
|
||||
"Tetrate Agent Router Service",
|
||||
"Enterprise router for AI models",
|
||||
TETRATE_DEFAULT_MODEL,
|
||||
TETRATE_KNOWN_MODELS.to_vec(),
|
||||
TETRATE_DOC_URL,
|
||||
vec![
|
||||
ConfigKey::new("TETRATE_API_KEY", true, true, None),
|
||||
ConfigKey::new(
|
||||
"TETRATE_HOST",
|
||||
false,
|
||||
false,
|
||||
Some("https://api.router.tetrate.ai"),
|
||||
),
|
||||
],
|
||||
)
|
||||
}
|
||||
|
||||
fn get_model_config(&self) -> ModelConfig {
|
||||
self.model.clone()
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
skip(self, system, messages, tools),
|
||||
fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
|
||||
)]
|
||||
async fn complete(
|
||||
&self,
|
||||
system: &str,
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
||||
// Create the base payload
|
||||
let payload = create_request_based_on_model(self, system, messages, tools)?;
|
||||
|
||||
// Make request
|
||||
let response = self
|
||||
.with_retry(|| async {
|
||||
let payload_clone = payload.clone();
|
||||
self.post(&payload_clone).await
|
||||
})
|
||||
.await?;
|
||||
|
||||
// Parse response
|
||||
let message = response_to_message(&response)?;
|
||||
let usage = response.get("usage").map(get_usage).unwrap_or_else(|| {
|
||||
tracing::debug!("Failed to get usage data");
|
||||
Usage::default()
|
||||
});
|
||||
let model = get_model(&response);
|
||||
emit_debug_trace(&self.model, &payload, &response, &usage);
|
||||
Ok((message, ProviderUsage::new(model, usage)))
|
||||
}
|
||||
|
||||
/// Fetch supported models from Tetrate Agent Router Service API (only models with tool support)
|
||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
||||
// Use the existing api_client which already has authentication configured
|
||||
let response = match self.api_client.response_get("v1/models").await {
|
||||
Ok(response) => response,
|
||||
Err(e) => {
|
||||
tracing::warn!("Failed to fetch models from Tetrate Agent Router Service API: {}, falling back to manual model entry", e);
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
|
||||
// Handle JSON parsing failures gracefully
|
||||
let json: serde_json::Value = match response.json().await {
|
||||
Ok(json) => json,
|
||||
Err(e) => {
|
||||
tracing::warn!("Failed to parse Tetrate Agent Router Service API response as JSON: {}, falling back to manual model entry", e);
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
|
||||
// Check for error in response
|
||||
if let Some(err_obj) = json.get("error") {
|
||||
let msg = err_obj
|
||||
.get("message")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("unknown error");
|
||||
tracing::warn!(
|
||||
"Tetrate Agent Router Service API returned an error: {}",
|
||||
msg
|
||||
);
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
// The response format from /v1/models is expected to be OpenAI-compatible
|
||||
// It should have a "data" field with an array of model objects
|
||||
let data = json.get("data").and_then(|v| v.as_array()).ok_or_else(|| {
|
||||
ProviderError::UsageError("Missing data field in JSON response".into())
|
||||
})?;
|
||||
|
||||
let mut models: Vec<String> = data
|
||||
.iter()
|
||||
.filter_map(|model| {
|
||||
// Get the model ID
|
||||
let id = model.get("id").and_then(|v| v.as_str())?;
|
||||
|
||||
// Check if the model supports computer_use (which indicates tool/function support)
|
||||
// The Tetrate API uses "supports_computer_use" instead of "supported_parameters"
|
||||
let supports_computer_use = model
|
||||
.get("supports_computer_use")
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
|
||||
if supports_computer_use {
|
||||
Some(id.to_string())
|
||||
} else {
|
||||
tracing::debug!(
|
||||
"Model '{}' does not support computer_use (tool support), skipping",
|
||||
id
|
||||
);
|
||||
None
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
// If no models with tool support were found, fall back to manual entry
|
||||
if models.is_empty() {
|
||||
tracing::warn!("No models with tool support found in Tetrate Agent Router Service API response, falling back to manual model entry");
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
models.sort();
|
||||
Ok(Some(models))
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user