diff --git a/app/src/pages/onboarding/steps/__tests__/ApiKeysStep.test.tsx b/app/src/pages/onboarding/steps/__tests__/ApiKeysStep.test.tsx index 7dd9063ae..31fb03dc3 100644 --- a/app/src/pages/onboarding/steps/__tests__/ApiKeysStep.test.tsx +++ b/app/src/pages/onboarding/steps/__tests__/ApiKeysStep.test.tsx @@ -1,7 +1,11 @@ import { fireEvent, screen, waitFor } from '@testing-library/react'; import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { setCloudProviderKey } from '../../../../services/api/aiSettingsApi'; +import { callCoreRpc } from '../../../../services/coreRpcClient'; import { renderWithProviders } from '../../../../test/test-utils'; +import { openUrl } from '../../../../utils/openUrl'; +import { isTauri } from '../../../../utils/tauriCommands/common'; import ApiKeysStep from '../ApiKeysStep'; vi.mock('../../../../services/coreRpcClient', () => ({ callCoreRpc: vi.fn() })); @@ -17,10 +21,12 @@ vi.mock('../../../../services/api/aiSettingsApi', () => ({ describe('ApiKeysStep OpenAI OAuth', () => { beforeEach(() => { vi.clearAllMocks(); + vi.mocked(isTauri).mockReturnValue(true); + vi.mocked(openUrl).mockResolvedValue(undefined); + vi.mocked(setCloudProviderKey).mockResolvedValue(undefined); }); it('shows connected badge when oauth status reports connected', async () => { - const { callCoreRpc } = await import('../../../../services/coreRpcClient'); vi.mocked(callCoreRpc).mockResolvedValueOnce({ result: { connected: true } }); renderWithProviders(); @@ -30,7 +36,6 @@ describe('ApiKeysStep OpenAI OAuth', () => { }); it('starts oauth and accepts pasted callback URL', async () => { - const { callCoreRpc } = await import('../../../../services/coreRpcClient'); vi.mocked(callCoreRpc) .mockResolvedValueOnce({ result: { connected: false } }) .mockResolvedValueOnce({ @@ -42,8 +47,6 @@ describe('ApiKeysStep OpenAI OAuth', () => { }) .mockResolvedValueOnce({ result: { connected: true } }); - const { openUrl } = await import('../../../../utils/openUrl'); - renderWithProviders(); fireEvent.click(await screen.findByTestId('onboarding-openai-oauth-connect')); @@ -71,4 +74,94 @@ describe('ApiKeysStep OpenAI OAuth', () => { expect(await screen.findByTestId('onboarding-openai-oauth-connected')).toBeInTheDocument(); }); + + it('shows a desktop-only error without calling core outside Tauri', async () => { + vi.mocked(isTauri).mockReturnValue(false); + + renderWithProviders(); + + fireEvent.click(screen.getByTestId('onboarding-openai-oauth-connect')); + + expect( + await screen.findByText('ChatGPT sign-in is only available in the desktop app.') + ).toBeInTheDocument(); + expect(callCoreRpc).not.toHaveBeenCalled(); + expect(openUrl).not.toHaveBeenCalled(); + }); + + it('reports an oauth start failure when core omits authUrl', async () => { + vi.mocked(callCoreRpc) + .mockResolvedValueOnce({ result: { connected: false } }) + .mockResolvedValueOnce({ result: { authUrl: ' ' } }); + + renderWithProviders(); + + fireEvent.click(await screen.findByTestId('onboarding-openai-oauth-connect')); + + expect( + await screen.findByText('Could not start ChatGPT sign-in. Try again or use an API key.') + ).toBeInTheDocument(); + expect(openUrl).not.toHaveBeenCalled(); + }); + + it('requires a pasted callback before completing oauth', async () => { + vi.mocked(callCoreRpc) + .mockResolvedValueOnce({ result: { connected: false } }) + .mockResolvedValueOnce({ + result: { authUrl: 'https://auth.openai.com/oauth/authorize?client_id=test' }, + }); + + renderWithProviders(); + + fireEvent.click(await screen.findByTestId('onboarding-openai-oauth-connect')); + await screen.findByTestId('onboarding-openai-oauth-callback-input'); + fireEvent.click(screen.getByTestId('onboarding-openai-oauth-complete')); + + expect( + await screen.findByText('Paste the redirect URL from your browser after signing in.') + ).toBeInTheDocument(); + expect(callCoreRpc).not.toHaveBeenCalledWith( + expect.objectContaining({ method: 'openhuman.inference_openai_oauth_complete' }) + ); + }); + + it('reports an oauth completion failure and keeps the callback form visible', async () => { + vi.mocked(callCoreRpc) + .mockResolvedValueOnce({ result: { connected: false } }) + .mockResolvedValueOnce({ + result: { authUrl: 'https://auth.openai.com/oauth/authorize?client_id=test' }, + }) + .mockRejectedValueOnce(new Error('state mismatch')); + + renderWithProviders(); + + fireEvent.click(await screen.findByTestId('onboarding-openai-oauth-connect')); + const input = await screen.findByTestId('onboarding-openai-oauth-callback-input'); + fireEvent.change(input, { + target: { value: 'http://127.0.0.1:1455/auth/callback?code=abc&state=wrong' }, + }); + fireEvent.click(screen.getByTestId('onboarding-openai-oauth-complete')); + + expect( + await screen.findByText( + 'ChatGPT sign-in did not complete. Check the redirect URL and try again.' + ) + ).toBeInTheDocument(); + expect(screen.getByTestId('onboarding-openai-oauth-callback-input')).toBeInTheDocument(); + }); + + it('continues without saving API keys when oauth is already connected', async () => { + const onNext = vi.fn(); + vi.mocked(callCoreRpc).mockResolvedValueOnce({ result: { connected: true } }); + + renderWithProviders(); + + await screen.findByTestId('onboarding-openai-oauth-connected'); + fireEvent.click(screen.getByTestId('onboarding-next-button')); + + await waitFor(() => { + expect(onNext).toHaveBeenCalledTimes(1); + }); + expect(setCloudProviderKey).not.toHaveBeenCalled(); + }); }); diff --git a/src/openhuman/inference/openai_oauth/flow.rs b/src/openhuman/inference/openai_oauth/flow.rs index e7bf88de4..cebabb2f5 100644 --- a/src/openhuman/inference/openai_oauth/flow.rs +++ b/src/openhuman/inference/openai_oauth/flow.rs @@ -233,7 +233,7 @@ pub fn disconnect_openai_oauth(config: &Config) -> Result Config { - let mut config = Config::default(); - config.config_path = tmp.path().join("config.toml"); - config + Config { + config_path: tmp.path().join("config.toml"), + ..Config::default() + } +} + +fn runtime() -> tokio::runtime::Runtime { + tokio::runtime::Runtime::new().unwrap() +} + +fn unsigned_jwt(payload: serde_json::Value) -> String { + use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _}; + + let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"none"}"#); + let payload = URL_SAFE_NO_PAD.encode(payload.to_string()); + format!("{header}.{payload}.") +} + +fn test_oauth_config(token_url: &'static str) -> OAuthConfig { + OAuthConfig { + client_id: "client-id", + client_secret: Some("client-secret"), + auth_url: "https://auth.example.test/oauth/authorize", + token_url, + scopes: &["scope-a", "scope-b"], + redirect_port: Some(1455), + callback_path: "/auth/callback", + redirect_uri_host: "127.0.0.1", + token_body: TokenBodyFormat::Form, + extra_auth_params: &[("prompt", "consent")], + state_strategy: StateStrategy::Random, + } } #[test] @@ -30,6 +65,41 @@ fn start_openai_oauth_returns_authorize_url() { assert!(!openai_oauth_status(&config).unwrap().connected); } +#[test] +fn build_authorize_url_includes_codex_pkce_and_extra_params() { + let url = build_authorize_url( + &test_oauth_config("https://token.example.test/oauth/token"), + "challenge-123", + "state-123", + "http://127.0.0.1:1455/auth/callback", + ); + let parsed = reqwest::Url::parse(&url).unwrap(); + let pairs = parsed + .query_pairs() + .into_owned() + .collect::>(); + + assert_eq!( + pairs.get("client_id").map(String::as_str), + Some("client-id") + ); + assert_eq!(pairs.get("response_type").map(String::as_str), Some("code")); + assert_eq!( + pairs.get("scope").map(String::as_str), + Some("scope-a scope-b") + ); + assert_eq!(pairs.get("state").map(String::as_str), Some("state-123")); + assert_eq!( + pairs.get("code_challenge").map(String::as_str), + Some("challenge-123") + ); + assert_eq!( + pairs.get("code_challenge_method").map(String::as_str), + Some("S256") + ); + assert_eq!(pairs.get("prompt").map(String::as_str), Some("consent")); +} + #[test] fn parse_callback_input_accepts_full_redirect_url() { let url = "http://127.0.0.1:1455/auth/callback?code=abc&state=xyz"; @@ -38,12 +108,71 @@ fn parse_callback_input_accepts_full_redirect_url() { assert_eq!(state, "xyz"); } +#[test] +fn parse_callback_input_accepts_raw_query_string() { + let (code, state) = parse_callback_input("code=abc%20123&state=xyz").unwrap(); + assert_eq!(code, "abc 123"); + assert_eq!(state, "xyz"); +} + #[test] fn parse_callback_input_rejects_missing_code() { let err = parse_callback_input("http://127.0.0.1:1455/auth/callback?state=xyz").unwrap_err(); assert!(err.contains("code")); } +#[test] +fn parse_callback_input_rejects_blank_invalid_and_missing_state() { + let blank = parse_callback_input(" ").unwrap_err(); + assert!(blank.contains("required")); + + let invalid = parse_callback_input("not-a-callback").unwrap_err(); + assert!(invalid.contains("invalid")); + + let missing_state = + parse_callback_input("http://127.0.0.1:1455/auth/callback?code=abc").unwrap_err(); + assert!(missing_state.contains("state")); +} + +#[test] +fn complete_openai_oauth_rejects_missing_pending_session() { + let tmp = tempdir().unwrap(); + let config = test_config(&tmp); + let err = runtime() + .block_on(complete_openai_oauth( + &config, + "http://127.0.0.1:1455/auth/callback?code=fake&state=state", + )) + .unwrap_err(); + assert!(err.contains("no pending OAuth session")); +} + +#[test] +fn complete_openai_oauth_rejects_expired_pending_session() { + let tmp = tempdir().unwrap(); + let config = test_config(&tmp); + std::fs::write( + tmp.path().join("openai-oauth-pending.json"), + serde_json::json!({ + "state": "state", + "verifier": "verifier", + "redirect_uri": "http://127.0.0.1:1455/auth/callback", + "created_at": 1_u64, + }) + .to_string(), + ) + .unwrap(); + + let err = runtime() + .block_on(complete_openai_oauth( + &config, + "http://127.0.0.1:1455/auth/callback?code=fake&state=state", + )) + .unwrap_err(); + assert!(err.contains("no pending OAuth session")); + assert!(!tmp.path().join("openai-oauth-pending.json").exists()); +} + #[test] fn complete_openai_oauth_rejects_state_mismatch() { let tmp = tempdir().unwrap(); @@ -53,13 +182,134 @@ fn complete_openai_oauth_rejects_state_mismatch() { "http://127.0.0.1:1455/auth/callback?code=fake&state=not-{}", start.state ); - let rt = tokio::runtime::Runtime::new().unwrap(); - let err = rt + let err = runtime() .block_on(complete_openai_oauth(&config, &callback)) .unwrap_err(); assert!(err.contains("state mismatch")); } +#[tokio::test] +async fn exchange_authorization_code_parses_successful_token_response() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "access_token": "access-token", + "refresh_token": "refresh-token", + "id_token": "id-token", + "expires_in": 3600, + }))) + .mount(&server) + .await; + let token_url: &'static str = Box::leak(format!("{}/token", server.uri()).into_boxed_str()); + + let token = exchange_authorization_code( + &test_oauth_config(token_url), + "code-123", + "state-123", + "verifier-123", + "http://127.0.0.1:1455/auth/callback", + ) + .await + .unwrap(); + + assert_eq!(token.access_token, "access-token"); + assert_eq!(token.refresh_token, "refresh-token"); + assert_eq!(token.id_token.as_deref(), Some("id-token")); + assert_eq!(token.expires_in, 3600); + assert!(token.issued_at > 0); +} + +#[tokio::test] +async fn exchange_authorization_code_reports_http_errors() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(400).set_body_string("bad auth code")) + .mount(&server) + .await; + let token_url: &'static str = Box::leak(format!("{}/token", server.uri()).into_boxed_str()); + + let err = exchange_authorization_code( + &test_oauth_config(token_url), + "code-123", + "state-123", + "verifier-123", + "http://127.0.0.1:1455/auth/callback", + ) + .await + .unwrap_err(); + + assert!(err.contains("HTTP 400")); + assert!(err.contains("bad auth code")); +} + +#[test] +fn persist_openai_oauth_token_stores_oauth_profile_with_metadata() { + let tmp = tempdir().unwrap(); + let config = test_config(&tmp); + let access_token = unsigned_jwt(serde_json::json!({ "sub": "acct_123" })); + let token = motosan_ai_oauth::Token { + access_token: access_token.clone(), + refresh_token: "refresh-token".into(), + id_token: Some("id-token".into()), + expires_in: 3600, + issued_at: 123, + }; + + let profile = persist_openai_oauth_token(&config, &token).unwrap(); + assert_eq!(profile.kind, AuthProfileKind::OAuth); + assert_eq!( + profile.metadata.get("account_id").map(String::as_str), + Some("acct_123") + ); + assert_eq!( + profile + .token_set + .as_ref() + .map(|set| set.access_token.as_str()), + Some(access_token.as_str()) + ); + assert_eq!( + profile + .token_set + .as_ref() + .and_then(|set| set.refresh_token.as_deref()), + Some("refresh-token") + ); + assert!(profile + .token_set + .as_ref() + .and_then(|set| set.expires_at) + .is_some()); + + let data = AuthProfilesStore::new(tmp.path(), false).load().unwrap(); + let stored = data.profiles.get(&profile.id).unwrap(); + assert_eq!(stored.id, profile.id); +} + +#[test] +fn openai_oauth_status_reports_token_profile_as_disconnected() { + let tmp = tempdir().unwrap(); + let config = test_config(&tmp); + let store = AuthProfilesStore::new(tmp.path(), false); + store + .upsert_profile( + AuthProfile::new_token( + OPENAI_PROVIDER_KEY, + OPENAI_OAUTH_PROFILE_NAME, + "sk-token-profile".to_string(), + ), + true, + ) + .unwrap(); + + let status = openai_oauth_status(&config).unwrap(); + assert!(!status.connected); + assert_eq!(status.auth_method.as_deref(), Some("token")); + assert!(status.profile_id.is_some()); +} + #[test] fn lookup_openai_bearer_token_prefers_api_key_over_oauth() { let tmp = tempdir().unwrap(); @@ -88,6 +338,105 @@ fn lookup_openai_bearer_token_prefers_api_key_over_oauth() { assert_eq!(token.as_deref(), Some("sk-api-key")); } +#[test] +fn lookup_openai_bearer_token_uses_oauth_when_api_key_missing() { + let tmp = tempdir().unwrap(); + let config = test_config(&tmp); + let store = AuthProfilesStore::new(tmp.path(), false); + let oauth_profile = AuthProfile::new_oauth( + OPENAI_PROVIDER_KEY, + OPENAI_OAUTH_PROFILE_NAME, + TokenSet { + access_token: "oauth-access".into(), + refresh_token: Some("refresh".into()), + id_token: None, + expires_at: Some(Utc::now() + Duration::hours(1)), + token_type: Some("Bearer".into()), + scope: None, + }, + ); + store.upsert_profile(oauth_profile, true).unwrap(); + + let token = lookup_openai_bearer_token(&config).unwrap(); + assert_eq!(token.as_deref(), Some("oauth-access")); +} + +#[test] +fn lookup_openai_bearer_token_uses_legacy_api_key_when_new_style_is_empty() { + let tmp = tempdir().unwrap(); + let config = test_config(&tmp); + let store = AuthProfilesStore::new(tmp.path(), false); + let oauth_profile = AuthProfile::new_oauth( + OPENAI_PROVIDER_KEY, + OPENAI_OAUTH_PROFILE_NAME, + TokenSet { + access_token: " ".into(), + refresh_token: None, + id_token: None, + expires_at: Some(Utc::now() + Duration::hours(1)), + token_type: Some("Bearer".into()), + scope: None, + }, + ); + store.upsert_profile(oauth_profile, true).unwrap(); + store + .upsert_profile( + AuthProfile::new_token("openai", "default", "sk-legacy-key".to_string()), + true, + ) + .unwrap(); + + let token = lookup_openai_bearer_token(&config).unwrap(); + assert_eq!(token.as_deref(), Some("sk-legacy-key")); +} + +#[test] +fn lookup_openai_bearer_token_keeps_expired_token_when_refresh_fails_without_runtime() { + let tmp = tempdir().unwrap(); + let config = test_config(&tmp); + let store = AuthProfilesStore::new(tmp.path(), false); + let oauth_profile = AuthProfile::new_oauth( + OPENAI_PROVIDER_KEY, + OPENAI_OAUTH_PROFILE_NAME, + TokenSet { + access_token: "expired-access".into(), + refresh_token: Some("refresh".into()), + id_token: None, + expires_at: Some(Utc::now() - Duration::minutes(5)), + token_type: Some("Bearer".into()), + scope: None, + }, + ); + store.upsert_profile(oauth_profile, true).unwrap(); + + let token = lookup_openai_bearer_token(&config).unwrap(); + assert_eq!(token.as_deref(), Some("expired-access")); +} + +#[test] +fn lookup_openai_bearer_token_returns_none_without_profiles_or_access_token() { + let tmp = tempdir().unwrap(); + let config = test_config(&tmp); + assert_eq!(lookup_openai_bearer_token(&config).unwrap(), None); + + let store = AuthProfilesStore::new(tmp.path(), false); + let empty_oauth_profile = AuthProfile::new_oauth( + OPENAI_PROVIDER_KEY, + OPENAI_OAUTH_PROFILE_NAME, + TokenSet { + access_token: " ".into(), + refresh_token: None, + id_token: None, + expires_at: Some(Utc::now() - Duration::hours(1)), + token_type: Some("Bearer".into()), + scope: None, + }, + ); + store.upsert_profile(empty_oauth_profile, true).unwrap(); + + assert_eq!(lookup_openai_bearer_token(&config).unwrap(), None); +} + #[test] fn disconnect_openai_oauth_clears_profile() { let tmp = tempdir().unwrap(); diff --git a/src/openhuman/inference/ops_tests.rs b/src/openhuman/inference/ops_tests.rs index 7d76274b1..771ac9d95 100644 --- a/src/openhuman/inference/ops_tests.rs +++ b/src/openhuman/inference/ops_tests.rs @@ -1,11 +1,16 @@ use super::*; +use crate::openhuman::credentials::profiles::{AuthProfile, AuthProfilesStore, TokenSet}; +use crate::openhuman::inference::openai_oauth::{OPENAI_OAUTH_PROFILE_NAME, OPENAI_PROVIDER_KEY}; +use chrono::{Duration, Utc}; use tempfile::tempdir; fn disabled_config() -> (Config, tempfile::TempDir) { let tmp = tempdir().expect("tempdir"); - let mut config = Config::default(); - config.workspace_dir = tmp.path().join("workspace"); - config.config_path = tmp.path().join("config.toml"); + let mut config = Config { + workspace_dir: tmp.path().join("workspace"), + config_path: tmp.path().join("config.toml"), + ..Config::default() + }; config.local_ai.runtime_enabled = false; config.local_ai.opt_in_confirmed = false; (config, tmp) @@ -109,3 +114,99 @@ async fn inference_presets_returns_recommended_tier() { assert!(outcome.value.get("recommended_tier").is_some()); assert!(outcome.value.get("presets").is_some()); } + +#[tokio::test] +async fn inference_openai_oauth_start_returns_authorize_payload() { + let (config, _tmp) = disabled_config(); + + let outcome = inference_openai_oauth_start(&config) + .await + .expect("oauth start"); + + assert!(outcome.value["authUrl"] + .as_str() + .unwrap() + .contains("auth.openai.com")); + assert_eq!( + outcome.value["redirectUri"].as_str(), + Some("http://127.0.0.1:1455/auth/callback") + ); + assert_eq!(outcome.logs, vec!["openai oauth authorize url ready"]); +} + +#[tokio::test] +async fn inference_openai_oauth_complete_surfaces_state_errors() { + let (config, _tmp) = disabled_config(); + let start = inference_openai_oauth_start(&config) + .await + .expect("oauth start"); + let state = start.value["state"].as_str().unwrap(); + let callback = format!("http://127.0.0.1:1455/auth/callback?code=fake&state=wrong-{state}"); + + let err = inference_openai_oauth_complete(&config, &callback) + .await + .expect_err("state mismatch should fail"); + + assert!(err.contains("state mismatch")); +} + +#[tokio::test] +async fn inference_openai_oauth_status_returns_connected_payload() { + let (config, tmp) = disabled_config(); + let store = AuthProfilesStore::new(tmp.path(), false); + store + .upsert_profile( + AuthProfile::new_oauth( + OPENAI_PROVIDER_KEY, + OPENAI_OAUTH_PROFILE_NAME, + TokenSet { + access_token: "oauth-access".into(), + refresh_token: None, + id_token: None, + expires_at: Some(Utc::now() + Duration::hours(1)), + token_type: Some("Bearer".into()), + scope: None, + }, + ), + true, + ) + .unwrap(); + + let outcome = inference_openai_oauth_status(&config) + .await + .expect("oauth status"); + + assert_eq!(outcome.value["connected"], true); + assert_eq!(outcome.value["authMethod"], "oauth"); + assert_eq!(outcome.logs, vec!["openai oauth status"]); +} + +#[tokio::test] +async fn inference_openai_oauth_disconnect_returns_removed_flag() { + let (config, tmp) = disabled_config(); + let store = AuthProfilesStore::new(tmp.path(), false); + store + .upsert_profile( + AuthProfile::new_oauth( + OPENAI_PROVIDER_KEY, + OPENAI_OAUTH_PROFILE_NAME, + TokenSet { + access_token: "oauth-access".into(), + refresh_token: None, + id_token: None, + expires_at: None, + token_type: Some("Bearer".into()), + scope: None, + }, + ), + true, + ) + .unwrap(); + + let outcome = inference_openai_oauth_disconnect(&config) + .await + .expect("oauth disconnect"); + + assert_eq!(outcome.value["disconnected"], true); + assert_eq!(outcome.logs, vec!["openai oauth disconnected"]); +} diff --git a/src/openhuman/inference/provider/factory_test.rs b/src/openhuman/inference/provider/factory_test.rs index 61df1314c..e57bdd51e 100644 --- a/src/openhuman/inference/provider/factory_test.rs +++ b/src/openhuman/inference/provider/factory_test.rs @@ -454,3 +454,22 @@ fn verify_session_active_called_for_custom_provider_not_for_openhuman() { "verify_session_active must reject config without session", ); } + +#[test] +fn lookup_key_for_slug_routes_openai_oauth_lookup_path() { + let tmp = TempDir::new().expect("tempdir"); + let config = config_in_tempdir(&tmp); + let auth = AuthService::new(tmp.path(), config.secrets.encrypt); + auth.store_provider_token( + "provider:openai", + "default", + "sk-openai", + Default::default(), + true, + ) + .expect("store openai token"); + + let token = lookup_key_for_slug("openai", &config).expect("lookup openai token"); + + assert_eq!(token, "sk-openai"); +} diff --git a/src/openhuman/inference/schemas.rs b/src/openhuman/inference/schemas.rs index 2109d5e91..70f70b9f6 100644 --- a/src/openhuman/inference/schemas.rs +++ b/src/openhuman/inference/schemas.rs @@ -701,9 +701,7 @@ fn handle_inference_openai_oauth_status(_params: Map) -> Controll fn handle_inference_openai_oauth_disconnect(_params: Map) -> ControllerFuture { Box::pin(async move { let config = config_rpc::load_config_with_timeout().await?; - to_json( - crate::openhuman::inference::rpc::inference_openai_oauth_disconnect(&config).await?, - ) + to_json(crate::openhuman::inference::rpc::inference_openai_oauth_disconnect(&config).await?) }) } diff --git a/src/openhuman/inference/schemas_tests.rs b/src/openhuman/inference/schemas_tests.rs index d1581d605..6085cc588 100644 --- a/src/openhuman/inference/schemas_tests.rs +++ b/src/openhuman/inference/schemas_tests.rs @@ -36,6 +36,10 @@ fn inference_schema_function_names_are_stable() { assert!(functions.contains(&"presets")); assert!(functions.contains(&"apply_preset")); assert!(functions.contains(&"diagnostics")); + assert!(functions.contains(&"openai_oauth_start")); + assert!(functions.contains(&"openai_oauth_complete")); + assert!(functions.contains(&"openai_oauth_status")); + assert!(functions.contains(&"openai_oauth_disconnect")); assert!(functions.contains(&"prompt")); assert!(functions.contains(&"vision_prompt")); assert!(functions.contains(&"embed")); @@ -64,6 +68,45 @@ fn inference_chat_schema_requires_messages() { .any(|field| field.name == "messages" && field.required)); } +#[test] +fn inference_openai_oauth_schemas_are_registered_with_expected_shapes() { + let registered: Vec<&str> = all_registered_controllers() + .into_iter() + .map(|controller| controller.schema.function) + .collect(); + for function in [ + "openai_oauth_start", + "openai_oauth_complete", + "openai_oauth_status", + "openai_oauth_disconnect", + ] { + assert!(registered.contains(&function), "missing {function}"); + let schema = schemas(function); + assert_eq!(schema.namespace, "inference"); + assert_eq!(schema.function, function); + assert!(!schema.description.is_empty()); + assert!(!schema.outputs.is_empty()); + } + + let complete = schemas("openai_oauth_complete"); + assert_eq!(complete.inputs.len(), 1); + assert_eq!(complete.inputs[0].name, "callback_url"); + assert!(complete.inputs[0].required); + + assert!(schemas("openai_oauth_start").inputs.is_empty()); + assert!(schemas("openai_oauth_status").inputs.is_empty()); + assert!(schemas("openai_oauth_disconnect").inputs.is_empty()); +} + +#[tokio::test] +async fn inference_openai_oauth_complete_handler_rejects_invalid_params() { + let params = Map::from_iter([("callback_url".to_string(), Value::Bool(true))]); + let err = handle_inference_openai_oauth_complete(params) + .await + .expect_err("invalid params"); + assert!(err.contains("invalid params")); +} + #[test] fn inference_unknown_schema_panics() { let panic = std::panic::catch_unwind(|| schemas("no_such_function"));