diff --git a/rust/src/core/http.rs b/rust/src/core/http.rs index db445f7ac6..7ad38b3ecf 100644 --- a/rust/src/core/http.rs +++ b/rust/src/core/http.rs @@ -15,7 +15,7 @@ pub fn credentialed_http_client_builder() -> reqwest::ClientBuilder { return attempt.follow(); }; - if is_same_origin_redirect(last_url, attempt.url()) { + if is_same_origin(last_url, attempt.url()) { attempt.follow() } else { attempt.stop() @@ -23,7 +23,7 @@ pub fn credentialed_http_client_builder() -> reqwest::ClientBuilder { })) } -fn is_same_origin_redirect(from: &Url, to: &Url) -> bool { +pub(crate) fn is_same_origin(from: &Url, to: &Url) -> bool { from.scheme() == to.scheme() && from.host_str() == to.host_str() && from.port_or_known_default() == to.port_or_known_default() @@ -39,7 +39,7 @@ mod tests { #[test] fn same_origin_redirect_allows_path_changes() { - assert!(is_same_origin_redirect( + assert!(is_same_origin( &url("https://example.com/a"), &url("https://example.com/b?x=1"), )); @@ -47,7 +47,7 @@ mod tests { #[test] fn same_origin_redirect_rejects_host_changes() { - assert!(!is_same_origin_redirect( + assert!(!is_same_origin( &url("https://example.com/a"), &url("https://evil.example/b"), )); @@ -55,7 +55,7 @@ mod tests { #[test] fn same_origin_redirect_rejects_scheme_changes() { - assert!(!is_same_origin_redirect( + assert!(!is_same_origin( &url("https://example.com/a"), &url("http://example.com/b"), )); diff --git a/rust/src/providers/mod.rs b/rust/src/providers/mod.rs index fb73ebd393..75e063f787 100755 --- a/rust/src/providers/mod.rs +++ b/rust/src/providers/mod.rs @@ -121,13 +121,26 @@ pub use zed::ZedProvider; pub(crate) fn browser_cookie_header( domains: &[&str], ) -> Result { - crate::browser::cookies::get_cookie_header_for_domains(domains).map_err(|error| match error { + crate::browser::cookies::get_cookie_header_for_domains(domains) + .map_err(map_browser_cookie_error) +} + +pub(crate) fn browser_cookies_for_domain( + domain: &str, +) -> Result, crate::core::ProviderError> { + crate::browser::cookies::get_cookies_for_domain(domain).map_err(map_browser_cookie_error) +} + +fn map_browser_cookie_error( + error: crate::browser::cookies::CookieError, +) -> crate::core::ProviderError { + match error { crate::browser::cookies::CookieError::BrowserNotInstalled | crate::browser::cookies::CookieError::NotFound(_) => { crate::core::ProviderError::NoCookies } _ => crate::core::ProviderError::Other(format!("Failed to read browser cookies: {error}")), - }) + } } pub(crate) fn resolve_api_key( diff --git a/rust/src/providers/ollama/mod.rs b/rust/src/providers/ollama/mod.rs index 3f94ec3c8a..8ea5047c92 100755 --- a/rust/src/providers/ollama/mod.rs +++ b/rust/src/providers/ollama/mod.rs @@ -9,6 +9,7 @@ use regex_lite::Regex; use reqwest::Url; use serde::Deserialize; +use crate::browser::cookies::{Cookie, CookieExtractor}; use crate::core::{ FetchContext, Provider, ProviderError, ProviderFetchResult, ProviderId, ProviderMetadata, RateWindow, SourceMode, UsageSnapshot, @@ -18,8 +19,18 @@ use crate::settings::ApiKeys; /// Ollama settings page URL const OLLAMA_SETTINGS_URL: &str = "https://ollama.com/settings"; const OLLAMA_TAGS_URL: &str = "https://ollama.com/api/tags"; +const OLLAMA_VALIDATION_URL: &str = "https://ollama.com/api/web_search"; const OLLAMA_COOKIE_DOMAIN: &str = "ollama.com"; const OLLAMA_SESSION_COOKIE_NAME: &str = "__Secure-session"; +const OLLAMA_SESSION_COOKIE_NAMES: &[&str] = &[ + "session", + OLLAMA_SESSION_COOKIE_NAME, + "ollama_session", + "__Host-ollama_session", + "wos-session", + "__Secure-next-auth.session-token", + "next-auth.session-token", +]; /// Ollama provider pub struct OllamaProvider { @@ -34,6 +45,20 @@ struct UsageBlock { reset_description: Option, } +enum OllamaCookieSource { + Manual(String), + Browser(Vec), +} + +impl OllamaCookieSource { + fn header_for_url(&self, url: &Url) -> Option { + match self { + Self::Manual(header) => should_attach_ollama_cookie(url).then(|| header.clone()), + Self::Browser(cookies) => ollama_cookie_header_for_url(cookies, url), + } + } +} + impl OllamaProvider { pub fn new() -> Self { Self { @@ -54,89 +79,16 @@ impl OllamaProvider { /// Fetch usage by scraping ollama.com/settings async fn fetch_usage_web(&self, ctx: &FetchContext) -> Result { - let cookie_header = self.resolve_cookie_header(ctx)?; + let cookies = self.resolve_cookie_source(ctx)?; let client = crate::core::credentialed_http_client_builder() .timeout(std::time::Duration::from_secs(ctx.web_timeout)) .redirect(reqwest::redirect::Policy::none()) .build() .map_err(|e| ProviderError::Other(e.to_string()))?; - - let mut current_url = + let start_url = Url::parse(OLLAMA_SETTINGS_URL).map_err(|e| ProviderError::Other(e.to_string()))?; - let mut resp = None; - - for _ in 0..5 { - let mut request = client - .get(current_url.clone()) - .header( - "Accept", - "text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8", - ) - .header( - "User-Agent", - "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/143.0.0.0 Safari/537.36", - ); - - if should_attach_ollama_cookie(¤t_url) { - request = request.header("Cookie", &cookie_header); - } - - let response = request.send().await?; - if response.status().is_redirection() { - let Some(location) = response.headers().get(reqwest::header::LOCATION) else { - return Err(ProviderError::Other( - "Ollama redirect missing Location header".to_string(), - )); - }; - let location = location - .to_str() - .map_err(|e| ProviderError::Other(e.to_string()))?; - let next_url = current_url - .join(location) - .map_err(|e| ProviderError::Other(e.to_string()))?; - if is_ollama_login_url(&next_url) { - return Err(ProviderError::AuthRequired); - } - if !should_attach_ollama_cookie(&next_url) { - return Err(ProviderError::AuthRequired); - } - current_url = next_url; - continue; - } - resp = Some(response); - break; - } - - let Some(resp) = resp else { - return Err(ProviderError::Other( - "Ollama returned too many redirects".to_string(), - )); - }; - - if resp.status() == reqwest::StatusCode::UNAUTHORIZED - || resp.status() == reqwest::StatusCode::FORBIDDEN - { - return Err(ProviderError::AuthRequired); - } - - // Check for redirect to login page - if is_ollama_login_url(resp.url()) { - return Err(ProviderError::AuthRequired); - } - - if !resp.status().is_success() { - return Err(ProviderError::Other(format!( - "Ollama returned status {}", - resp.status() - ))); - } - - let html = resp - .text() - .await - .map_err(|e| ProviderError::Other(e.to_string()))?; - + let html = fetch_settings_html_at(&client, &cookies, start_url).await?; self.parse_usage_html(&html) } @@ -146,9 +98,52 @@ impl OllamaProvider { .timeout(std::time::Duration::from_secs(ctx.web_timeout.max(1))) .build() .map_err(|e| ProviderError::Other(e.to_string()))?; + + let validation_url = + Url::parse(OLLAMA_VALIDATION_URL).map_err(|e| ProviderError::Other(e.to_string()))?; + let tags_url = + Url::parse(OLLAMA_TAGS_URL).map_err(|e| ProviderError::Other(e.to_string()))?; + Self::fetch_usage_api_at(&client, &api_key, validation_url, tags_url).await + } + + async fn fetch_usage_api_at( + client: &reqwest::Client, + api_key: &str, + validation_url: Url, + tags_url: Url, + ) -> Result { + let api_key = clean_secret(Some(api_key)).ok_or(ProviderError::AuthRequired)?; + if !crate::core::is_same_origin(&validation_url, &tags_url) { + return Err(ProviderError::Other( + "Ollama API endpoints must share an origin.".to_string(), + )); + } + + let validation = client + .post(validation_url) + .bearer_auth(&api_key) + .header("Accept", "application/json") + .header("Content-Type", "application/json") + .header("User-Agent", "CodexBar/1.0") + .body(r#"{"query":""}"#) + .send() + .await?; + match validation.status() { + reqwest::StatusCode::OK | reqwest::StatusCode::BAD_REQUEST => {} + reqwest::StatusCode::UNAUTHORIZED | reqwest::StatusCode::FORBIDDEN => { + return Err(ollama_api_key_error()); + } + status => { + return Err(ProviderError::Other(format!( + "Ollama API validation returned status {}", + status.as_u16() + ))); + } + } + let response = client - .get(OLLAMA_TAGS_URL) - .bearer_auth(api_key) + .get(tags_url) + .bearer_auth(&api_key) .header("Accept", "application/json") .header("User-Agent", "CodexBar/1.0") .send() @@ -239,37 +234,34 @@ impl OllamaProvider { } } - /// Resolve cookie header from manual cookies, browser import, or context - fn resolve_cookie_header(&self, ctx: &FetchContext) -> Result { + /// Resolve cookies from manual cookies, browser import, or context. + fn resolve_cookie_source( + &self, + ctx: &FetchContext, + ) -> Result { // Check manual cookie header first if let Some(ref cookie) = ctx.manual_cookie_header && let Some(header) = Self::normalize_cookie_header(cookie) { - return Ok(header); + return has_recognized_ollama_session_cookie(&header) + .then_some(OllamaCookieSource::Manual(header)) + .ok_or(ProviderError::NoCookies); } // Try browser cookie extraction - match crate::providers::browser_cookie_header(&[OLLAMA_COOKIE_DOMAIN]) { - Ok(header) if !header.is_empty() => { - // Validate that we have a recognized session cookie - const SESSION_COOKIE_NAMES: &[&str] = &[ - "session", - "__Secure-session", - "ollama_session", - "__Host-ollama_session", - "__Secure-next-auth.session-token", - "next-auth.session-token", - ]; - let has_session = SESSION_COOKIE_NAMES - .iter() - .any(|name| header.contains(name)); - if has_session { - Ok(header) - } else { - Err(ProviderError::NoCookies) - } + match crate::providers::browser_cookies_for_domain(OLLAMA_COOKIE_DOMAIN) { + Ok(cookies) => { + let source = OllamaCookieSource::Browser(cookies); + source + .header_for_url( + &Url::parse(OLLAMA_SETTINGS_URL) + .map_err(|e| ProviderError::Other(e.to_string()))?, + ) + .is_some() + .then_some(source) + .ok_or(ProviderError::NoCookies) } - Ok(_) | Err(ProviderError::NoCookies) => Err(ProviderError::NoCookies), + Err(ProviderError::NoCookies) => Err(ProviderError::NoCookies), Err(err) => Err(err), } } @@ -490,9 +482,138 @@ fn should_attach_ollama_cookie(url: &Url) -> bool { .is_some_and(|host| host.eq_ignore_ascii_case(OLLAMA_COOKIE_DOMAIN)) } -fn is_ollama_login_url(url: &Url) -> bool { +fn has_recognized_ollama_session_cookie(header: &str) -> bool { + header.split(';').any(|pair| { + let name = pair.trim().split_once('=').map(|(name, _)| name.trim()); + name.is_some_and(is_recognized_ollama_session_cookie_name) + }) +} + +fn ollama_cookie_header_for_url(cookies: &[Cookie], url: &Url) -> Option { + let cookies: Vec<_> = cookies + .iter() + .filter(|cookie| cookie_applies_to_ollama_url(cookie, url)) + .cloned() + .collect(); + let header = CookieExtractor::build_cookie_header(&cookies); + has_recognized_ollama_session_cookie(&header).then_some(header) +} + +fn cookie_applies_to_ollama_url(cookie: &Cookie, url: &Url) -> bool { + let domain = cookie + .domain + .trim() + .trim_end_matches('.') + .to_ascii_lowercase(); + let path = if cookie.path.is_empty() { + "/" + } else { + cookie.path.as_str() + }; + let request_path = url.path(); + should_attach_ollama_cookie(url) + && (domain == OLLAMA_COOKIE_DOMAIN + || domain.strip_prefix('.') == Some(OLLAMA_COOKIE_DOMAIN)) + && (path == "/" + || request_path == path + || (request_path.starts_with(path) + && (path.ends_with('/') || request_path.as_bytes().get(path.len()) == Some(&b'/')))) +} + +fn is_recognized_ollama_session_cookie_name(name: &str) -> bool { + OLLAMA_SESSION_COOKIE_NAMES.contains(&name) + || is_chunked_nextauth_cookie_name(name, "__Secure-next-auth.session-token") + || is_chunked_nextauth_cookie_name(name, "next-auth.session-token") +} + +fn is_chunked_nextauth_cookie_name(name: &str, base_name: &str) -> bool { + name.strip_prefix(base_name) + .and_then(|suffix| suffix.strip_prefix('.')) + .is_some_and(|suffix| { + !suffix.is_empty() && suffix.bytes().all(|byte| byte.is_ascii_digit()) + }) +} + +fn is_ollama_sign_in_redirect(url: &Url) -> bool { + if url.scheme() != "https" { + return false; + } + let Some(host) = url.host_str().map(str::to_ascii_lowercase) else { + return false; + }; let path = url.path().to_ascii_lowercase(); - path.contains("/login") || path.contains("/signin") + if host == OLLAMA_COOKIE_DOMAIN || host == "www.ollama.com" { + return path == "/signin" || path.starts_with("/signin/") || path.contains("/login"); + } + host == "signin.ollama.com" + || (host.ends_with(".workos.com") && path.starts_with("/user_management/authorize")) +} + +async fn fetch_settings_html_at( + client: &reqwest::Client, + source: &OllamaCookieSource, + start_url: Url, +) -> Result { + let mut current_url = start_url; + + for _ in 0..5 { + let mut request = client + .get(current_url.clone()) + .header( + "Accept", + "text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8", + ) + .header( + "User-Agent", + "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/143.0.0.0 Safari/537.36", + ); + if let Some(cookie_header) = source.header_for_url(¤t_url) { + request = request.header("Cookie", cookie_header); + } + + let response = request.send().await?; + if response.status().is_redirection() { + let Some(location) = response.headers().get(reqwest::header::LOCATION) else { + return Err(ProviderError::Other( + "Ollama redirect missing Location header".to_string(), + )); + }; + let location = location + .to_str() + .map_err(|e| ProviderError::Other(e.to_string()))?; + let next_url = current_url + .join(location) + .map_err(|e| ProviderError::Other(e.to_string()))?; + if is_ollama_sign_in_redirect(&next_url) + || !crate::core::is_same_origin(¤t_url, &next_url) + { + return Err(ProviderError::AuthRequired); + } + current_url = next_url; + continue; + } + + if response.status() == reqwest::StatusCode::UNAUTHORIZED + || response.status() == reqwest::StatusCode::FORBIDDEN + || is_ollama_sign_in_redirect(response.url()) + { + return Err(ProviderError::AuthRequired); + } + if !response.status().is_success() { + return Err(ProviderError::Other(format!( + "Ollama returned status {}", + response.status() + ))); + } + return response + .text() + .await + .map_err(|e| ProviderError::Other(e.to_string())); + } + + Err(ProviderError::Other( + "Ollama returned too many redirects".to_string(), + )) } fn ollama_api_key_error() -> ProviderError { @@ -533,6 +654,83 @@ mod tests { assert_eq!(OllamaProvider::normalize_cookie_header("Cookie: "), None); } + #[test] + fn recognizes_exact_authkit_and_nextauth_session_cookie_names() { + assert!(has_recognized_ollama_session_cookie( + "wos-session=auth; theme=dark" + )); + assert!(has_recognized_ollama_session_cookie( + "__Secure-next-auth.session-token.0=auth" + )); + assert!(!has_recognized_ollama_session_cookie( + "notwos-session=auth; theme=dark" + )); + assert!(!has_recognized_ollama_session_cookie( + "next-auth.session-token.evil=auth" + )); + assert!(!has_recognized_ollama_session_cookie("theme=dark")); + } + + #[test] + fn limits_browser_cookie_headers_to_ollama_settings_scope() { + use crate::browser::cookies::Cookie; + + let cookie = |name: &str, domain: &str, path: &str| Cookie { + name: name.to_string(), + value: "test".to_string(), + domain: domain.to_string(), + path: path.to_string(), + expires: None, + is_secure: true, + is_http_only: true, + }; + let cookies = [ + cookie("wos-session", ".ollama.com", "/"), + cookie("wos-session", "signin.ollama.com", "/"), + cookie("__Secure-session", "ollama.com", "/signin"), + ]; + + assert_eq!( + ollama_cookie_header_for_url( + &cookies, + &Url::parse("https://ollama.com/settings").unwrap() + ) + .as_deref(), + Some("wos-session=test") + ); + assert_eq!( + ollama_cookie_header_for_url( + &[cookie("__Secure-session", "ollama.com", "/settings")], + &Url::parse("https://ollama.com/api/tags").unwrap() + ), + None + ); + assert_eq!( + ollama_cookie_header_for_url( + &[cookie("__Secure-session", "ollama.com", "/settings")], + &Url::parse("https://ollama.com/settings/account").unwrap() + ) + .as_deref(), + Some("__Secure-session=test") + ); + let source = OllamaCookieSource::Browser(vec![ + cookie("__Secure-session", "ollama.com", "/settings"), + cookie("wos-session", "ollama.com", "/api"), + ]); + assert_eq!( + source + .header_for_url(&Url::parse("https://ollama.com/settings").unwrap()) + .as_deref(), + Some("__Secure-session=test") + ); + assert_eq!( + source + .header_for_url(&Url::parse("https://ollama.com/api/models").unwrap()) + .as_deref(), + Some("wos-session=test") + ); + } + #[test] fn only_attaches_web_cookie_to_https_ollama_urls() { assert!(should_attach_ollama_cookie( @@ -546,6 +744,181 @@ mod tests { )); } + #[test] + fn recognizes_workos_signin_redirects_as_expired_sessions() { + assert!(is_ollama_sign_in_redirect( + &Url::parse("https://signin.ollama.com/?client_id=test").unwrap() + )); + assert!(is_ollama_sign_in_redirect( + &Url::parse("https://auth.workos.com/user_management/authorize?client_id=test") + .unwrap() + )); + assert!(!is_ollama_sign_in_redirect( + &Url::parse("https://auth.workos.com/other").unwrap() + )); + assert!(!is_ollama_sign_in_redirect( + &Url::parse("http://signin.ollama.com/").unwrap() + )); + } + + #[tokio::test] + async fn settings_fetch_follows_same_origin_redirects() { + let mut server = mockito::Server::new_async().await; + let first = server + .mock("GET", "/settings") + .with_status(302) + .with_header("location", "/settings/account") + .create_async() + .await; + let second = server + .mock("GET", "/settings/account") + .with_status(200) + .with_body("usage") + .create_async() + .await; + let client = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .unwrap(); + + let html = fetch_settings_html_at( + &client, + &OllamaCookieSource::Manual("__Secure-session=test".to_string()), + Url::parse(&format!("{}/settings", server.url())).unwrap(), + ) + .await + .unwrap(); + + first.assert_async().await; + second.assert_async().await; + assert_eq!(html, "usage"); + } + + #[tokio::test] + async fn settings_fetch_stops_before_following_signin_or_workos_redirects() { + for location in [ + "https://signin.ollama.com/?client_id=test", + "https://auth.workos.com/user_management/authorize?client_id=test", + ] { + let mut server = mockito::Server::new_async().await; + let first = server + .mock("GET", "/settings") + .with_status(302) + .with_header("location", location) + .create_async() + .await; + let client = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .unwrap(); + + let error = fetch_settings_html_at( + &client, + &OllamaCookieSource::Manual("__Secure-session=test".to_string()), + Url::parse(&format!("{}/settings", server.url())).unwrap(), + ) + .await + .unwrap_err(); + + first.assert_async().await; + assert!(matches!(error, ProviderError::AuthRequired)); + } + } + + #[tokio::test] + async fn settings_fetch_reports_redirect_exhaustion() { + let mut server = mockito::Server::new_async().await; + let redirect = server + .mock("GET", "/settings") + .expect(5) + .with_status(302) + .with_header("location", "/settings") + .create_async() + .await; + let client = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .unwrap(); + + let error = fetch_settings_html_at( + &client, + &OllamaCookieSource::Manual("__Secure-session=test".to_string()), + Url::parse(&format!("{}/settings", server.url())).unwrap(), + ) + .await + .unwrap_err(); + + redirect.assert_async().await; + assert_eq!(error.to_string(), "Ollama returned too many redirects"); + } + + #[tokio::test] + async fn validates_trimmed_key_before_fetching_public_model_catalog() { + let mut server = mockito::Server::new_async().await; + let validation = server + .mock("POST", "/api/web_search") + .match_header("authorization", "Bearer ollama-key") + .match_header("content-type", "application/json") + .match_body(r#"{"query":""}"#) + .with_status(400) + .create_async() + .await; + let catalog = server + .mock("GET", "/api/tags") + .match_header("authorization", "Bearer ollama-key") + .with_status(200) + .with_body(r#"{"models":[{"name":"gpt-oss"}]}"#) + .create_async() + .await; + let client = reqwest::Client::new(); + + let snapshot = OllamaProvider::fetch_usage_api_at( + &client, + " ollama-key ", + Url::parse(&format!("{}/api/web_search", server.url())).unwrap(), + Url::parse(&format!("{}/api/tags", server.url())).unwrap(), + ) + .await + .unwrap(); + + validation.assert_async().await; + catalog.assert_async().await; + assert_eq!(snapshot.login_method.as_deref(), Some("API key")); + } + + #[tokio::test] + async fn rejects_unproven_validation_responses_before_catalog_fetch() { + let mut server = mockito::Server::new_async().await; + let validation = server + .mock("POST", "/api/web_search") + .with_status(422) + .create_async() + .await; + let catalog = server + .mock("GET", "/api/tags") + .expect(0) + .with_status(200) + .create_async() + .await; + let client = reqwest::Client::new(); + + let error = OllamaProvider::fetch_usage_api_at( + &client, + "ollama-key", + Url::parse(&format!("{}/api/web_search", server.url())).unwrap(), + Url::parse(&format!("{}/api/tags", server.url())).unwrap(), + ) + .await + .unwrap_err(); + + validation.assert_async().await; + catalog.assert_async().await; + assert_eq!( + error.to_string(), + "Ollama API validation returned status 422" + ); + } + #[test] fn strips_wrapping_quotes_from_api_key() { assert_eq!(