Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
000559bc9d
|
||
|
|
3aede58a11
|
||
|
|
cab1e72b97
|
||
|
|
cd4bf245e9
|
||
|
|
82bf6176f8
|
||
|
|
d407eb5a6a
|
+8
-2
@@ -320,6 +320,14 @@
|
||||
# - https://ai.google.dev/api/rest/v1beta/models/streamGenerateContent
|
||||
- provider: gemini
|
||||
models:
|
||||
- name: gemini-3.6-flash
|
||||
max_input_tokens: 1048576
|
||||
max_output_tokens: 65536
|
||||
input_price: 1.5
|
||||
output_price: 7.5
|
||||
supports_function_calling: true
|
||||
reasoning_levels: [minimal, low, medium, high]
|
||||
default_reasoning_level: medium
|
||||
- name: gemini-3.5-flash
|
||||
max_input_tokens: 1048576
|
||||
max_output_tokens: 65536
|
||||
@@ -360,8 +368,6 @@
|
||||
output_price: 0
|
||||
supports_vision: true
|
||||
supports_function_calling: true
|
||||
reasoning_levels: [low, medium, high]
|
||||
default_reasoning_effort: low
|
||||
- name: gemini-2.5-pro
|
||||
max_input_tokens: 1048576
|
||||
max_output_tokens: 65536
|
||||
|
||||
+148
-19
@@ -2,7 +2,7 @@ use super::access_token::{is_valid_access_token, set_access_token};
|
||||
use super::openai_compatible_oauth::OpenAICompatibleOAuthProvider;
|
||||
use super::{ClientConfig, ProviderModels};
|
||||
use crate::config::paths;
|
||||
use anyhow::{Result, anyhow, bail};
|
||||
use anyhow::{Context, Result, anyhow, bail};
|
||||
use base64::Engine;
|
||||
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
|
||||
use chrono::Utc;
|
||||
@@ -248,20 +248,24 @@ async fn run_pkce_flow(provider: &dyn OAuthProvider, client_name: &str) -> Resul
|
||||
let _ = open::that(&authorize_url);
|
||||
|
||||
let (code, returned_state) = if use_callback_listener {
|
||||
listen_for_oauth_callback(&redirect_uri)?
|
||||
let (code, state) = listen_for_oauth_callback(&redirect_uri)?;
|
||||
(code, Some(state))
|
||||
} else {
|
||||
let input = Text::new("Paste the authorization code:").prompt()?;
|
||||
let parts: Vec<&str> = input.splitn(2, '#').collect();
|
||||
if parts.len() != 2 {
|
||||
bail!("Invalid authorization code format. Expected format: <code>#<state>");
|
||||
}
|
||||
(parts[0].to_string(), parts[1].to_string())
|
||||
let input = Text::new("Paste the authorization code or callback URL:").prompt()?;
|
||||
parse_paste_input(input.trim())?
|
||||
};
|
||||
|
||||
if returned_state != state {
|
||||
bail!(
|
||||
"OAuth state mismatch: expected '{state}', got '{returned_state}'. \
|
||||
This may indicate a CSRF attack or a stale authorization attempt."
|
||||
if let Some(returned) = returned_state.as_deref() {
|
||||
if returned != state {
|
||||
bail!(
|
||||
"OAuth state mismatch: expected '{state}', got '{returned}'. \
|
||||
This may indicate a CSRF attack or a stale authorization attempt."
|
||||
);
|
||||
}
|
||||
} else {
|
||||
eprintln!(
|
||||
"Warning: no state returned in the paste; skipping CSRF check. \
|
||||
If your provider's callback page shows a URL or code#state string, paste that instead."
|
||||
);
|
||||
}
|
||||
|
||||
@@ -390,7 +394,10 @@ async fn run_device_code_flow(provider: &dyn OAuthProvider, client_name: &str) -
|
||||
}
|
||||
let form: HashMap<&str, &str> = device_params.iter().copied().collect();
|
||||
|
||||
let mut device_request = client.post(device_auth_url).form(&form);
|
||||
let mut device_request = client
|
||||
.post(device_auth_url)
|
||||
.header("Accept", "application/json")
|
||||
.form(&form);
|
||||
for (key, value) in provider.extra_token_headers() {
|
||||
device_request = device_request.header(key, value);
|
||||
}
|
||||
@@ -468,6 +475,7 @@ async fn run_device_code_flow(provider: &dyn OAuthProvider, client_name: &str) -
|
||||
token_params.push(("code_verifier", verifier.as_str()));
|
||||
}
|
||||
let token_response: Value = build_token_request(&client, provider, &token_params)
|
||||
.header("Accept", "application/json")
|
||||
.send()
|
||||
.await?
|
||||
.json()
|
||||
@@ -475,10 +483,10 @@ async fn run_device_code_flow(provider: &dyn OAuthProvider, client_name: &str) -
|
||||
|
||||
if let Some(access_token) = token_response["access_token"].as_str() {
|
||||
let refresh_token = token_response["refresh_token"].as_str().map(str::to_string);
|
||||
let expires_in_secs = token_response["expires_in"].as_i64().ok_or_else(|| {
|
||||
anyhow!("Missing expires_in in device_code token response: {token_response}")
|
||||
})?;
|
||||
let expires_at = Utc::now().timestamp() + expires_in_secs;
|
||||
let expires_at = match token_response["expires_in"].as_i64() {
|
||||
Some(secs) => Utc::now().timestamp() + secs,
|
||||
None => i64::MAX,
|
||||
};
|
||||
let account_id = provider.extract_account_id(&token_response);
|
||||
|
||||
let tokens = OAuthTokens {
|
||||
@@ -670,6 +678,35 @@ fn build_token_request(
|
||||
request
|
||||
}
|
||||
|
||||
fn parse_paste_input(input: &str) -> Result<(String, Option<String>)> {
|
||||
if input.is_empty() {
|
||||
bail!("Empty input; paste the code, code#state, or callback URL from your browser.");
|
||||
}
|
||||
|
||||
if input.starts_with("http://") || input.starts_with("https://") {
|
||||
let parsed =
|
||||
Url::parse(input).with_context(|| format!("Failed to parse pasted URL: {input}"))?;
|
||||
let code = parsed
|
||||
.query_pairs()
|
||||
.find(|(k, _)| k == "code")
|
||||
.map(|(_, v)| v.to_string())
|
||||
.ok_or_else(|| {
|
||||
anyhow!("Pasted URL is missing the ?code= parameter. Paste the URL you were redirected to after approving.")
|
||||
})?;
|
||||
let state = parsed
|
||||
.query_pairs()
|
||||
.find(|(k, _)| k == "state")
|
||||
.map(|(_, v)| v.to_string());
|
||||
return Ok((code, state));
|
||||
}
|
||||
|
||||
if let Some((code, state)) = input.split_once('#') {
|
||||
return Ok((code.to_string(), Some(state.to_string())));
|
||||
}
|
||||
|
||||
Ok((input.to_string(), None))
|
||||
}
|
||||
|
||||
fn listen_for_oauth_callback(redirect_uri: &str) -> Result<(String, String)> {
|
||||
let url: Url = redirect_uri.parse()?;
|
||||
let host = url.host_str().unwrap_or("127.0.0.1");
|
||||
@@ -1095,7 +1132,7 @@ echo_pkce_in_token_exchange: true
|
||||
#[test]
|
||||
fn openai_compatible_provider_prefers_redirect_uri_over_port() {
|
||||
let mut cfg = base_config();
|
||||
cfg.redirect_uri = Some("https://custom.example/cb".into());
|
||||
cfg.redirect_uri = Some("http://127.0.0.1:9000/cb".into());
|
||||
cfg.redirect_port = Some(9999);
|
||||
|
||||
let provider = OpenAICompatibleOAuthProvider {
|
||||
@@ -1105,7 +1142,7 @@ echo_pkce_in_token_exchange: true
|
||||
|
||||
assert_eq!(
|
||||
provider.fixed_redirect_uri().as_deref(),
|
||||
Some("https://custom.example/cb")
|
||||
Some("http://127.0.0.1:9000/cb")
|
||||
);
|
||||
}
|
||||
|
||||
@@ -1222,6 +1259,98 @@ echo_pkce_in_token_exchange: true
|
||||
assert!(provider.use_pkce_in_device_flow());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_paste_input_full_callback_url() {
|
||||
let (code, state) =
|
||||
parse_paste_input("https://provider.example/oauth/callback?code=abc123&state=xyz")
|
||||
.unwrap();
|
||||
assert_eq!(code, "abc123");
|
||||
assert_eq!(state.as_deref(), Some("xyz"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_paste_input_url_without_state() {
|
||||
let (code, state) =
|
||||
parse_paste_input("https://provider.example/oauth/callback?code=abc123").unwrap();
|
||||
assert_eq!(code, "abc123");
|
||||
assert!(state.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_paste_input_code_state_fragment() {
|
||||
let (code, state) = parse_paste_input("abc123#xyz").unwrap();
|
||||
assert_eq!(code, "abc123");
|
||||
assert_eq!(state.as_deref(), Some("xyz"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_paste_input_bare_code() {
|
||||
let (code, state) = parse_paste_input("abc123").unwrap();
|
||||
assert_eq!(code, "abc123");
|
||||
assert!(state.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_paste_input_url_missing_code_fails() {
|
||||
let err =
|
||||
parse_paste_input("https://provider.example/oauth/callback?state=xyz").unwrap_err();
|
||||
assert!(err.to_string().contains("code"), "unexpected error: {err}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_paste_input_empty_fails() {
|
||||
let err = parse_paste_input("").unwrap_err();
|
||||
assert!(err.to_string().contains("Empty"), "unexpected error: {err}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_compatible_provider_fixed_redirect_uri_none_for_public_url() {
|
||||
let mut cfg = base_config();
|
||||
cfg.redirect_uri = Some("https://provider.example.com/callback".into());
|
||||
cfg.redirect_port = None;
|
||||
|
||||
let provider = OpenAICompatibleOAuthProvider {
|
||||
config: cfg,
|
||||
client_name: "test".into(),
|
||||
};
|
||||
|
||||
assert!(provider.fixed_redirect_uri().is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_compatible_provider_fixed_redirect_uri_some_for_localhost() {
|
||||
let mut cfg = base_config();
|
||||
cfg.redirect_uri = Some("http://127.0.0.1:9999/cb".into());
|
||||
cfg.redirect_port = None;
|
||||
|
||||
let provider = OpenAICompatibleOAuthProvider {
|
||||
config: cfg,
|
||||
client_name: "test".into(),
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
provider.fixed_redirect_uri().as_deref(),
|
||||
Some("http://127.0.0.1:9999/cb")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_compatible_provider_fixed_redirect_uri_some_for_localhost_hostname() {
|
||||
let mut cfg = base_config();
|
||||
cfg.redirect_uri = Some("http://localhost:9999/cb".into());
|
||||
cfg.redirect_port = None;
|
||||
|
||||
let provider = OpenAICompatibleOAuthProvider {
|
||||
config: cfg,
|
||||
client_name: "test".into(),
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
provider.fixed_redirect_uri().as_deref(),
|
||||
Some("http://localhost:9999/cb")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oauth_config_serde_roundtrip_device_code_yaml() {
|
||||
let yaml = r#"
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
use url::Url;
|
||||
|
||||
use super::oauth::{OAuthConfig, OAuthFlow, OAuthProvider, TokenRequestFormat};
|
||||
|
||||
pub struct OpenAICompatibleOAuthProvider {
|
||||
@@ -5,6 +7,13 @@ pub struct OpenAICompatibleOAuthProvider {
|
||||
pub client_name: String,
|
||||
}
|
||||
|
||||
fn is_loopback_uri(uri: &str) -> bool {
|
||||
Url::parse(uri)
|
||||
.ok()
|
||||
.and_then(|u| u.host_str().map(str::to_string))
|
||||
.is_some_and(|host| matches!(host.as_str(), "127.0.0.1" | "localhost" | "[::1]" | "::1"))
|
||||
}
|
||||
|
||||
impl OAuthProvider for OpenAICompatibleOAuthProvider {
|
||||
fn provider_name(&self) -> &str {
|
||||
&self.client_name
|
||||
@@ -70,7 +79,11 @@ impl OAuthProvider for OpenAICompatibleOAuthProvider {
|
||||
|
||||
fn fixed_redirect_uri(&self) -> Option<String> {
|
||||
if let Some(uri) = &self.config.redirect_uri {
|
||||
return Some(uri.clone());
|
||||
return if is_loopback_uri(uri) {
|
||||
Some(uri.clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
}
|
||||
if let Some(port) = self.config.redirect_port {
|
||||
return Some(format!("http://127.0.0.1:{port}/callback"));
|
||||
|
||||
@@ -322,7 +322,11 @@ fn gemini_extract_chat_completions_text(data: &Value) -> Result<ChatCompletionsO
|
||||
bail!("Invalid response data: {data}");
|
||||
}
|
||||
}
|
||||
let output = ChatCompletionsOutput { text, tool_calls, ..Default::default() };
|
||||
let output = ChatCompletionsOutput {
|
||||
text,
|
||||
tool_calls,
|
||||
..Default::default()
|
||||
};
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
@@ -435,7 +439,7 @@ pub fn gemini_build_chat_completions_body(
|
||||
body["generationConfig"]["topP"] = v.into();
|
||||
}
|
||||
if let Some(v) = reasoning_effort {
|
||||
body["generation_config"]["thinking_level"] = v.into();
|
||||
body["generationConfig"]["thinking_config"] = json!({"thinking_level": v});
|
||||
}
|
||||
|
||||
if let Some(functions) = functions {
|
||||
|
||||
Reference in New Issue
Block a user