pr_01m47d15m3e54sn21z27rpy5n9/services/integrations/src/models.rs

134 lines5,645 bytesCodeBlame
1//! A workspace's own model providers: where each one's API is, how it takes
2//! its key, and checking that the key works. The requests themselves go
3//! through the model proxy, which never lets a sandbox see a key.
4
5use g1t_contracts::integrations::{ConnectionConfig, Provider};
6use worker::{Method, Result};
7
8use crate::http;
9
10/// Where requests go: without `/v1` for Anthropic's API, with the version
11/// for OpenAI's (`…/v1`, or Gemini's `…/v1beta/openai`).
12pub fn base_url(provider: Provider, config: &ConnectionConfig) -> String {
13 let given = config.base_url.as_deref().unwrap_or_default().trim_end_matches('/');
14 match provider {
15 Provider::AnthropicEndpoint => given.trim_end_matches("/v1").to_owned(),
16 Provider::OpenaiEndpoint => given.to_owned(),
17 // Azure's v1 API, on the workspace's own resource.
18 Provider::AzureOpenai => format!("{}/openai/v1", given.trim_end_matches("/openai/v1").trim_end_matches("/openai")),
19 _ => provider.spec().base_url.to_owned(),
20 }
21}
22
23/// The header the key goes in; `authorization` means `Bearer <key>`.
24pub fn auth_header(provider: Provider, config: &ConnectionConfig) -> String {
25 match provider {
26 Provider::AnthropicEndpoint | Provider::OpenaiEndpoint => config
27 .auth_header
28 .clone()
29 .unwrap_or_else(|| provider.spec().auth_header.to_owned()),
30 _ => provider.spec().auth_header.to_owned(),
31 }
32}
33
34/// Whether a model id is one an agent could use: not embeddings, images,
35/// speech or moderation.
36fn for_chat(id: &str) -> bool {
37 let id = id.to_ascii_lowercase();
38 !["embed", "tts", "whisper", "dall-e", "image", "moderation", "audio", "transcribe", "realtime", "search", "aqa", "imagen", "veo"]
39 .iter()
40 .any(|word| id.contains(word))
41}
42
43/// Asks the provider for its models with the key. `Ok` with what to say and
44/// the models it offers; `Err` with what went wrong.
45pub async fn test(
46 provider: Provider,
47 config: &ConnectionConfig,
48 key: Option<&str>,
49 gateway_token: Option<&str>,
50) -> Result<std::result::Result<(String, Vec<String>), String>> {
51 let base = base_url(provider, config);
52 let url = match provider.api() {
53 "anthropic" => format!("{base}/v1/models?limit=100"),
54 _ => format!("{base}/models"),
55 };
56 let header = auth_header(provider, config);
57 let bearer = key.map(|key| format!("Bearer {key}"));
58 let mut headers = vec![("anthropic-version", "2023-06-01")];
59 if let Some(key) = key {
60 if header == "authorization" {
61 headers.push(("authorization", bearer.as_deref().unwrap_or_default()));
62 } else {
63 headers.push((header.as_str(), key));
64 }
65 }
66 let gateway = gateway_token.map(|token| format!("Bearer {token}"));
67 if let Some(gateway) = gateway.as_deref() {
68 headers.push(("cf-aig-authorization", gateway));
69 }
70 let answer = http::send(Method::Get, &url, &headers, None).await?;
71 let system = provider.label();
72 if answer.ok() {
73 let mut models: Vec<String> = answer.json()["data"]
74 .as_array()
75 .map(|data| {
76 data.iter()
77 .filter_map(|model| model["id"].as_str())
78 // Gemini names models `models/gemini-…`.
79 .map(|id| id.trim_start_matches("models/").to_owned())
80 .filter(|id| for_chat(id))
81 .collect()
82 })
83 .unwrap_or_default();
84 models.sort();
85 models.truncate(1000);
86 let message = match models.len() {
87 0 => format!("{system} accepted the key."),
88 count => format!("{system} accepted the key and offers {count} models."),
89 };
90 return Ok(Ok((message, models)));
91 }
92 // A proxy may answer messages but not list models: it was reached, and
93 // whether the key works shows on the first run.
94 if answer.status == 404 && matches!(provider, Provider::AnthropicEndpoint | Provider::OpenaiEndpoint) {
95 return Ok(Ok((
96 "Reached the endpoint. It does not list models, so the key will be checked on the first run.".to_owned(),
97 Vec::new(),
98 )));
99 }
100 Ok(Err(answer.problem(system)))
101}
102
103#[cfg(test)]
104mod tests {
105 use super::*;
106
107 #[test]
108 fn each_provider_has_its_address_and_header() {
109 let config = ConnectionConfig {
110 base_url: Some("https://llm.acme.dev/v1/".to_owned()),
111 ..ConnectionConfig::default()
112 };
113 assert_eq!(base_url(Provider::Openai, &config), "https://api.openai.com/v1");
114 assert_eq!(base_url(Provider::OpenaiEndpoint, &config), "https://llm.acme.dev/v1");
115 assert_eq!(base_url(Provider::AnthropicEndpoint, &config), "https://llm.acme.dev");
116 assert_eq!(auth_header(Provider::Gemini, &config), "authorization");
117 assert_eq!(auth_header(Provider::Anthropic, &config), "x-api-key");
118 assert_eq!(auth_header(Provider::AzureOpenai, &config), "api-key");
119 assert_eq!(base_url(Provider::Groq, &config), "https://api.groq.com/openai/v1");
120 let azure = ConnectionConfig {
121 base_url: Some("https://acme.openai.azure.com/".to_owned()),
122 ..ConnectionConfig::default()
123 };
124 assert_eq!(base_url(Provider::AzureOpenai, &azure), "https://acme.openai.azure.com/openai/v1");
125 }
126
127 #[test]
128 fn only_models_that_can_chat_are_offered() {
129 assert!(for_chat("gpt-5"));
130 assert!(for_chat("gemini-2.5-pro"));
131 assert!(!for_chat("text-embedding-3-large"));
132 assert!(!for_chat("gpt-image-1"));
133 }
134}