Compare commits
6
Commits
6f2594712f
...
000559bc9d
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
000559bc9d
|
||
|
|
3aede58a11
|
||
|
|
cab1e72b97
|
||
|
|
cd4bf245e9
|
||
|
|
82bf6176f8
|
||
|
|
d407eb5a6a
|
+8
-2
@@ -320,6 +320,14 @@
|
|||||||
# - https://ai.google.dev/api/rest/v1beta/models/streamGenerateContent
|
# - https://ai.google.dev/api/rest/v1beta/models/streamGenerateContent
|
||||||
- provider: gemini
|
- provider: gemini
|
||||||
models:
|
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
|
- name: gemini-3.5-flash
|
||||||
max_input_tokens: 1048576
|
max_input_tokens: 1048576
|
||||||
max_output_tokens: 65536
|
max_output_tokens: 65536
|
||||||
@@ -360,8 +368,6 @@
|
|||||||
output_price: 0
|
output_price: 0
|
||||||
supports_vision: true
|
supports_vision: true
|
||||||
supports_function_calling: true
|
supports_function_calling: true
|
||||||
reasoning_levels: [low, medium, high]
|
|
||||||
default_reasoning_effort: low
|
|
||||||
- name: gemini-2.5-pro
|
- name: gemini-2.5-pro
|
||||||
max_input_tokens: 1048576
|
max_input_tokens: 1048576
|
||||||
max_output_tokens: 65536
|
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::openai_compatible_oauth::OpenAICompatibleOAuthProvider;
|
||||||
use super::{ClientConfig, ProviderModels};
|
use super::{ClientConfig, ProviderModels};
|
||||||
use crate::config::paths;
|
use crate::config::paths;
|
||||||
use anyhow::{Result, anyhow, bail};
|
use anyhow::{Context, Result, anyhow, bail};
|
||||||
use base64::Engine;
|
use base64::Engine;
|
||||||
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
|
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
|
||||||
use chrono::Utc;
|
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 _ = open::that(&authorize_url);
|
||||||
|
|
||||||
let (code, returned_state) = if use_callback_listener {
|
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 {
|
} else {
|
||||||
let input = Text::new("Paste the authorization code:").prompt()?;
|
let input = Text::new("Paste the authorization code or callback URL:").prompt()?;
|
||||||
let parts: Vec<&str> = input.splitn(2, '#').collect();
|
parse_paste_input(input.trim())?
|
||||||
if parts.len() != 2 {
|
|
||||||
bail!("Invalid authorization code format. Expected format: <code>#<state>");
|
|
||||||
}
|
|
||||||
(parts[0].to_string(), parts[1].to_string())
|
|
||||||
};
|
};
|
||||||
|
|
||||||
if returned_state != state {
|
if let Some(returned) = returned_state.as_deref() {
|
||||||
bail!(
|
if returned != state {
|
||||||
"OAuth state mismatch: expected '{state}', got '{returned_state}'. \
|
bail!(
|
||||||
This may indicate a CSRF attack or a stale authorization attempt."
|
"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 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() {
|
for (key, value) in provider.extra_token_headers() {
|
||||||
device_request = device_request.header(key, value);
|
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()));
|
token_params.push(("code_verifier", verifier.as_str()));
|
||||||
}
|
}
|
||||||
let token_response: Value = build_token_request(&client, provider, &token_params)
|
let token_response: Value = build_token_request(&client, provider, &token_params)
|
||||||
|
.header("Accept", "application/json")
|
||||||
.send()
|
.send()
|
||||||
.await?
|
.await?
|
||||||
.json()
|
.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() {
|
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 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(|| {
|
let expires_at = match token_response["expires_in"].as_i64() {
|
||||||
anyhow!("Missing expires_in in device_code token response: {token_response}")
|
Some(secs) => Utc::now().timestamp() + secs,
|
||||||
})?;
|
None => i64::MAX,
|
||||||
let expires_at = Utc::now().timestamp() + expires_in_secs;
|
};
|
||||||
let account_id = provider.extract_account_id(&token_response);
|
let account_id = provider.extract_account_id(&token_response);
|
||||||
|
|
||||||
let tokens = OAuthTokens {
|
let tokens = OAuthTokens {
|
||||||
@@ -670,6 +678,35 @@ fn build_token_request(
|
|||||||
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)> {
|
fn listen_for_oauth_callback(redirect_uri: &str) -> Result<(String, String)> {
|
||||||
let url: Url = redirect_uri.parse()?;
|
let url: Url = redirect_uri.parse()?;
|
||||||
let host = url.host_str().unwrap_or("127.0.0.1");
|
let host = url.host_str().unwrap_or("127.0.0.1");
|
||||||
@@ -1095,7 +1132,7 @@ echo_pkce_in_token_exchange: true
|
|||||||
#[test]
|
#[test]
|
||||||
fn openai_compatible_provider_prefers_redirect_uri_over_port() {
|
fn openai_compatible_provider_prefers_redirect_uri_over_port() {
|
||||||
let mut cfg = base_config();
|
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);
|
cfg.redirect_port = Some(9999);
|
||||||
|
|
||||||
let provider = OpenAICompatibleOAuthProvider {
|
let provider = OpenAICompatibleOAuthProvider {
|
||||||
@@ -1105,7 +1142,7 @@ echo_pkce_in_token_exchange: true
|
|||||||
|
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
provider.fixed_redirect_uri().as_deref(),
|
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());
|
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]
|
#[test]
|
||||||
fn oauth_config_serde_roundtrip_device_code_yaml() {
|
fn oauth_config_serde_roundtrip_device_code_yaml() {
|
||||||
let yaml = r#"
|
let yaml = r#"
|
||||||
|
|||||||
@@ -1,3 +1,5 @@
|
|||||||
|
use url::Url;
|
||||||
|
|
||||||
use super::oauth::{OAuthConfig, OAuthFlow, OAuthProvider, TokenRequestFormat};
|
use super::oauth::{OAuthConfig, OAuthFlow, OAuthProvider, TokenRequestFormat};
|
||||||
|
|
||||||
pub struct OpenAICompatibleOAuthProvider {
|
pub struct OpenAICompatibleOAuthProvider {
|
||||||
@@ -5,6 +7,13 @@ pub struct OpenAICompatibleOAuthProvider {
|
|||||||
pub client_name: String,
|
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 {
|
impl OAuthProvider for OpenAICompatibleOAuthProvider {
|
||||||
fn provider_name(&self) -> &str {
|
fn provider_name(&self) -> &str {
|
||||||
&self.client_name
|
&self.client_name
|
||||||
@@ -70,7 +79,11 @@ impl OAuthProvider for OpenAICompatibleOAuthProvider {
|
|||||||
|
|
||||||
fn fixed_redirect_uri(&self) -> Option<String> {
|
fn fixed_redirect_uri(&self) -> Option<String> {
|
||||||
if let Some(uri) = &self.config.redirect_uri {
|
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 {
|
if let Some(port) = self.config.redirect_port {
|
||||||
return Some(format!("http://127.0.0.1:{port}/callback"));
|
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}");
|
bail!("Invalid response data: {data}");
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
let output = ChatCompletionsOutput { text, tool_calls, ..Default::default() };
|
let output = ChatCompletionsOutput {
|
||||||
|
text,
|
||||||
|
tool_calls,
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
Ok(output)
|
Ok(output)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -435,7 +439,7 @@ pub fn gemini_build_chat_completions_body(
|
|||||||
body["generationConfig"]["topP"] = v.into();
|
body["generationConfig"]["topP"] = v.into();
|
||||||
}
|
}
|
||||||
if let Some(v) = reasoning_effort {
|
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 {
|
if let Some(functions) = functions {
|
||||||
|
|||||||
Reference in New Issue
Block a user