fix: Implement a CredentialStore for auth (#5741)
This commit is contained in:
@@ -1,9 +1,11 @@
|
||||
mod persist;
|
||||
|
||||
use axum::extract::{Query, State};
|
||||
use axum::response::Html;
|
||||
use axum::routing::get;
|
||||
use axum::Router;
|
||||
use minijinja::render;
|
||||
use rmcp::transport::auth::OAuthState;
|
||||
use rmcp::transport::auth::{CredentialStore, OAuthState, StoredCredentials};
|
||||
use rmcp::transport::AuthorizationManager;
|
||||
use serde::Deserialize;
|
||||
use std::net::SocketAddr;
|
||||
@@ -11,9 +13,7 @@ use std::sync::Arc;
|
||||
use tokio::sync::{oneshot, Mutex};
|
||||
use tracing::warn;
|
||||
|
||||
use crate::oauth::persist::{clear_credentials, load_cached_state, save_credentials};
|
||||
|
||||
mod persist;
|
||||
use crate::oauth::persist::GooseCredentialStore;
|
||||
|
||||
const CALLBACK_TEMPLATE: &str = include_str!("oauth_callback.html");
|
||||
|
||||
@@ -32,18 +32,21 @@ pub async fn oauth_flow(
|
||||
mcp_server_url: &String,
|
||||
name: &String,
|
||||
) -> Result<AuthorizationManager, anyhow::Error> {
|
||||
if let Ok(oauth_state) = load_cached_state(mcp_server_url, name).await {
|
||||
if let Some(authorization_manager) = oauth_state.into_authorization_manager() {
|
||||
if authorization_manager.refresh_token().await.is_ok() {
|
||||
return Ok(authorization_manager);
|
||||
}
|
||||
let credential_store = GooseCredentialStore::new(name.clone());
|
||||
let mut auth_manager = AuthorizationManager::new(mcp_server_url).await?;
|
||||
auth_manager.set_credential_store(credential_store.clone());
|
||||
|
||||
if auth_manager.initialize_from_store().await? {
|
||||
if auth_manager.refresh_token().await.is_ok() {
|
||||
return Ok(auth_manager);
|
||||
}
|
||||
|
||||
if let Err(e) = clear_credentials(name) {
|
||||
if let Err(e) = credential_store.clear().await {
|
||||
warn!("error clearing bad credentials: {}", e);
|
||||
}
|
||||
}
|
||||
|
||||
// No existing credentials or they were invalid - need to do the full oauth flow
|
||||
let (code_sender, code_receiver) = oneshot::channel::<CallbackParams>();
|
||||
let app_state = AppState {
|
||||
code_receiver: Arc::new(Mutex::new(Some(code_sender))),
|
||||
@@ -74,6 +77,7 @@ pub async fn oauth_flow(
|
||||
});
|
||||
|
||||
let mut oauth_state = OAuthState::new(mcp_server_url, None).await?;
|
||||
|
||||
let redirect_uri = format!("http://localhost:{}/oauth_callback", used_addr.port());
|
||||
oauth_state
|
||||
.start_authorization(&[], redirect_uri.as_str(), Some("goose"))
|
||||
@@ -91,13 +95,20 @@ pub async fn oauth_flow(
|
||||
} = code_receiver.await?;
|
||||
oauth_state.handle_callback(&auth_code, &csrf_token).await?;
|
||||
|
||||
if let Err(e) = save_credentials(name, &oauth_state).await {
|
||||
warn!("Failed to save credentials: {}", e);
|
||||
}
|
||||
let (client_id, token_response) = oauth_state.get_credentials().await?;
|
||||
|
||||
let auth_manager = oauth_state
|
||||
let mut auth_manager = oauth_state
|
||||
.into_authorization_manager()
|
||||
.ok_or_else(|| anyhow::anyhow!("Failed to get authorization manager"))?;
|
||||
|
||||
credential_store
|
||||
.save(StoredCredentials {
|
||||
client_id,
|
||||
token_response,
|
||||
})
|
||||
.await?;
|
||||
|
||||
auth_manager.set_credential_store(credential_store);
|
||||
|
||||
Ok(auth_manager)
|
||||
}
|
||||
|
||||
@@ -1,71 +1,54 @@
|
||||
use oauth2::{basic::BasicTokenType, EmptyExtraTokenFields, StandardTokenResponse};
|
||||
use reqwest::IntoUrl;
|
||||
use rmcp::transport::{auth::OAuthState, AuthError};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use rmcp::transport::auth::{AuthError, CredentialStore, StoredCredentials};
|
||||
|
||||
use crate::config::Config;
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SerializableCredentials {
|
||||
pub client_id: String,
|
||||
pub token_response: Option<StandardTokenResponse<EmptyExtraTokenFields, BasicTokenType>>,
|
||||
/// Goose-specific credential store that uses the Config system
|
||||
///
|
||||
/// This implementation stores OAuth credentials in the goose configuration
|
||||
/// system, which handles secure storage (e.g., keychain integration).
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct GooseCredentialStore {
|
||||
name: String,
|
||||
}
|
||||
|
||||
fn secret_key(name: &str) -> String {
|
||||
format!("oauth_creds_{name}")
|
||||
}
|
||||
impl GooseCredentialStore {
|
||||
pub fn new(name: String) -> Self {
|
||||
Self { name }
|
||||
}
|
||||
|
||||
pub async fn save_credentials(
|
||||
name: &str,
|
||||
oauth_state: &OAuthState,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let config = Config::global();
|
||||
let (client_id, token_response) = oauth_state.get_credentials().await?;
|
||||
|
||||
let credentials = SerializableCredentials {
|
||||
client_id,
|
||||
token_response,
|
||||
};
|
||||
|
||||
let key = secret_key(name);
|
||||
config.set_secret(&key, &credentials)?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn load_credentials(
|
||||
name: &str,
|
||||
) -> Result<SerializableCredentials, Box<dyn std::error::Error>> {
|
||||
let config = Config::global();
|
||||
let key = secret_key(name);
|
||||
let credentials: SerializableCredentials = config.get_secret(&key)?;
|
||||
|
||||
Ok(credentials)
|
||||
}
|
||||
|
||||
pub fn clear_credentials(name: &str) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let config = Config::global();
|
||||
|
||||
Ok(config.delete_secret(&secret_key(name))?)
|
||||
}
|
||||
|
||||
pub async fn load_cached_state<U: IntoUrl>(
|
||||
base_url: U,
|
||||
name: &str,
|
||||
) -> Result<OAuthState, AuthError> {
|
||||
let credentials = load_credentials(name)
|
||||
.await
|
||||
.map_err(|e| AuthError::InternalError(format!("Failed to load credentials: {}", e)))?;
|
||||
|
||||
if let Some(token_response) = credentials.token_response {
|
||||
let mut oauth_state = OAuthState::new(base_url, None).await?;
|
||||
oauth_state
|
||||
.set_credentials(&credentials.client_id, token_response)
|
||||
.await?;
|
||||
Ok(oauth_state)
|
||||
} else {
|
||||
Err(AuthError::InternalError(
|
||||
"No token response in cached credentials".to_string(),
|
||||
))
|
||||
fn secret_key(&self) -> String {
|
||||
format!("oauth_creds_{}", self.name)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl CredentialStore for GooseCredentialStore {
|
||||
async fn load(&self) -> Result<Option<StoredCredentials>, AuthError> {
|
||||
let config = Config::global();
|
||||
let key = self.secret_key();
|
||||
|
||||
match config.get_secret::<StoredCredentials>(&key) {
|
||||
Ok(credentials) => Ok(Some(credentials)),
|
||||
Err(_) => Ok(None), // No credentials found
|
||||
}
|
||||
}
|
||||
|
||||
async fn save(&self, credentials: StoredCredentials) -> Result<(), AuthError> {
|
||||
let config = Config::global();
|
||||
let key = self.secret_key();
|
||||
|
||||
config
|
||||
.set_secret(&key, &credentials)
|
||||
.map_err(|e| AuthError::InternalError(format!("Failed to save credentials: {}", e)))
|
||||
}
|
||||
|
||||
async fn clear(&self) -> Result<(), AuthError> {
|
||||
let config = Config::global();
|
||||
let key = self.secret_key();
|
||||
|
||||
config
|
||||
.delete_secret(&key)
|
||||
.map_err(|e| AuthError::InternalError(format!("Failed to clear credentials: {}", e)))
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user