Simplify enterprise OAuth login helpers and expand callback tests (#44616)

## What changed

- Consolidate callback address and port parsing in `enterprise_callback_settings`.
- Reuse `StoredOAuthTokens::has_refresh_token` for credential validation.
- Keep `EnterpriseOAuthHttpClient` private and inline removal of the `resource` parameter from authorization code token requests.
- Share the `complete_login` test helper across login and logout tests.

## Testing

Expand callback settings coverage to check default and configured ports, ports from callback URLs, and rejection of conflicting ports alongside loopback address validation.

GitOrigin-RevId: 42c9be5663ce13b7349cefac2c6a20c4fa974aef
This commit is contained in:
Nick Steele
2026-09-10 18:25:32 +00:00
committed by copyberry
parent 3422443ec4
commit 86661eb626
3 changed files with 81 additions and 91 deletions

View File

@@ -15,7 +15,6 @@ use codex_exec_server::HttpClient;
use http::Method;
use http::header::CONTENT_LENGTH;
use http::header::CONTENT_TYPE;
use oauth2::TokenResponse;
use rmcp::transport::AuthorizationManager;
use rmcp::transport::auth::AuthorizationMetadata;
use rmcp::transport::auth::OAuthHttpClient;
@@ -247,11 +246,19 @@ pub(crate) fn enterprise_callback_settings(
if client_id.is_none_or(|client_id| client_id.trim().is_empty()) {
bail!("enterprise IdP login requires its registered client ID");
}
let ip = enterprise_callback_bind_ip(callback_url)?;
let registered_port = callback_url
.map(Url::parse)
.transpose()?
.and_then(|url| url.port());
let (ip, registered_port) = if let Some(callback_url) = callback_url {
validate_ema_oauth_endpoint(callback_url, "enterprise IdP callback URL")?;
let callback = Url::parse(callback_url)?;
let ip = match (callback.scheme(), callback.host()) {
("http", Some(Host::Domain("localhost"))) => Ipv4Addr::LOCALHOST.into(),
("http", Some(Host::Ipv4(ip))) if ip.is_loopback() => ip.into(),
("http", Some(Host::Ipv6(ip))) if ip.is_loopback() => ip.into(),
_ => bail!("enterprise IdP callback URL must use an HTTP loopback address"),
};
(ip, callback.port())
} else {
(Ipv4Addr::LOCALHOST.into(), None)
};
if callback_port
.zip(registered_port)
.is_some_and(|(configured, registered)| configured != registered)
@@ -261,32 +268,13 @@ pub(crate) fn enterprise_callback_settings(
Ok((ip, callback_port.or(registered_port)))
}
fn enterprise_callback_bind_ip(callback_url: Option<&str>) -> Result<IpAddr> {
let Some(callback_url) = callback_url else {
return Ok(Ipv4Addr::LOCALHOST.into());
};
validate_ema_oauth_endpoint(callback_url, "enterprise IdP callback URL")?;
let callback = Url::parse(callback_url)?;
if callback.scheme() == "http" {
match callback.host() {
Some(Host::Domain("localhost")) => return Ok(Ipv4Addr::LOCALHOST.into()),
Some(Host::Ipv4(ip)) if ip.is_loopback() => return Ok(ip.into()),
Some(Host::Ipv6(ip)) if ip.is_loopback() => return Ok(ip.into()),
_ => {}
}
}
bail!("enterprise IdP callback URL must use an HTTP loopback address")
}
fn validate_enterprise_credentials(stored: &StoredOAuthTokens) -> Result<()> {
let credentials = &stored.token_response.0;
if credentials
.refresh_token()
.is_none_or(|refresh_token| refresh_token.secret().trim().is_empty())
{
if !stored.has_refresh_token() {
bail!("enterprise IdP login did not return a refresh token");
}
let assertion = credentials
let assertion = stored
.token_response
.0
.extra_fields()
.0
.get("id_token")
@@ -330,15 +318,9 @@ pub(crate) fn enterprise_authorization_url(auth_url: &str) -> Result<String> {
Ok(url.to_string())
}
pub(crate) fn without_oauth_resource(encoded: &[u8]) -> String {
url::form_urlencoded::Serializer::new(String::new())
.extend_pairs(url::form_urlencoded::parse(encoded).filter(|(key, _)| key != "resource"))
.finish()
}
/// rmcp supplies a resource indicator for MCP OAuth, but the independent OIDC
/// login must not request the IdP issuer as a protected-resource audience.
pub(crate) struct EnterpriseOAuthHttpClient(pub(crate) Arc<dyn OAuthHttpClient>);
struct EnterpriseOAuthHttpClient(Arc<dyn OAuthHttpClient>);
impl OAuthHttpClient for EnterpriseOAuthHttpClient {
fn execute(&self, mut request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> {
@@ -357,8 +339,13 @@ impl OAuthHttpClient for EnterpriseOAuthHttpClient {
&& url::form_urlencoded::parse(request.request.body())
.any(|(key, value)| key == "grant_type" && value == "authorization_code")
{
let body = without_oauth_resource(request.request.body()).into_bytes();
*request.request.body_mut() = body;
let body = url::form_urlencoded::Serializer::new(String::new())
.extend_pairs(
url::form_urlencoded::parse(request.request.body())
.filter(|(key, _)| key != "resource"),
)
.finish();
*request.request.body_mut() = body.into_bytes();
request.request.headers_mut().remove(CONTENT_LENGTH);
}
self.0.execute(request)

View File

@@ -16,6 +16,7 @@ use keyring::credential::CredentialApi;
use keyring::credential::CredentialBuilderApi;
use keyring::credential::CredentialPersistence;
use keyring::mock::MockCredential;
use oauth2::TokenResponse;
use pretty_assertions::assert_eq;
use serde_json::json;
use sha2::Digest;
@@ -75,6 +76,17 @@ async fn login(issuer: &str, callback_url: Option<&str>) -> Result<EnterpriseOAu
.await
}
async fn complete_login(issuer: &str) -> Result<EnterpriseOAuthCredentials> {
let handle = login(issuer, /*callback_url*/ None).await?;
callback(
&handle.authorization_url(),
issuer,
/*provider_error*/ false,
)
.await?;
handle.wait().await
}
async fn metadata(server: &MockServer, issuer: &str) {
Mock::given(method("GET")).and(path("/.well-known/oauth-authorization-server/idp"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
@@ -160,15 +172,30 @@ async fn enterprise_callback_errors_and_sdk_logs_exclude_credentials() -> Result
#[test]
fn enterprise_callback_requires_loopback() -> Result<()> {
for (callback, expected) in [
(None, "127.0.0.1"),
(Some("http://localhost/callback"), "127.0.0.1"),
(Some("http://127.0.0.2/callback"), "127.0.0.2"),
(Some("http://[::1]/callback"), "::1"),
let issuer = "https://idp.example";
let client_id = Some("enterprise-client");
for (callback, port, expected_ip, expected_port) in [
(None, None, "127.0.0.1", None),
(None, Some(8080), "127.0.0.1", Some(8080)),
(Some("http://localhost/callback"), None, "127.0.0.1", None),
(Some("http://127.0.0.2/callback"), None, "127.0.0.2", None),
(Some("http://[::1]/callback"), None, "::1", None),
(
Some("http://localhost:8080/callback"),
None,
"127.0.0.1",
Some(8080),
),
(
Some("http://localhost:8080/callback"),
Some(8080),
"127.0.0.1",
Some(8080),
),
] {
assert_eq!(
enterprise_callback_bind_ip(callback)?,
expected.parse::<IpAddr>()?
enterprise_callback_settings(issuer, client_id, callback, port)?,
(expected_ip.parse::<IpAddr>()?, expected_port),
);
}
for callback in [
@@ -177,8 +204,25 @@ fn enterprise_callback_requires_loopback() -> Result<()> {
"https://127.0.0.1/callback",
"http://remote.example/callback",
] {
assert!(enterprise_callback_bind_ip(Some(callback)).is_err());
assert!(
enterprise_callback_settings(
issuer,
client_id,
Some(callback),
/*callback_port*/ None
)
.is_err()
);
}
assert!(
enterprise_callback_settings(
issuer,
client_id,
Some("http://localhost:8080/callback"),
Some(9090),
)
.is_err()
);
Ok(())
}
@@ -317,14 +361,7 @@ async fn enterprise_public_api_storage_and_privacy() -> Result<()> {
// Cancellation while blocked on the actual credential lock cannot leave a
// detached persistence worker that writes after the other process releases it.
let canceled = login(&issuer, /*callback_url*/ None).await?;
callback(
&canceled.authorization_url(),
&issuer,
/*provider_error*/ false,
)
.await?;
let canceled = canceled.wait().await?;
let canceled = complete_login(&issuer).await?;
let guard = EnterpriseOAuthCredentialGuard::acquire(
CREDENTIAL_NAME,
&issuer,
@@ -345,23 +382,8 @@ async fn enterprise_public_api_storage_and_privacy() -> Result<()> {
);
// Rejected old attempts neither write nor delete a newer grant.
let old = login(&issuer, /*callback_url*/ None).await?;
callback(
&old.authorization_url(),
&issuer,
/*provider_error*/ false,
)
.await?;
let old = old.wait().await?;
let winner = login(&issuer, /*callback_url*/ None).await?;
callback(
&winner.authorization_url(),
&issuer,
/*provider_error*/ false,
)
.await?;
winner
.wait()
let old = complete_login(&issuer).await?;
complete_login(&issuer)
.await?
.commit_if(|| async { Some(()) })
.await?;
@@ -377,15 +399,7 @@ async fn enterprise_public_api_storage_and_privacy() -> Result<()> {
// Inject raw account identifiers into the actual keyring adapter's error chain.
keyring.fail.store(true, Ordering::SeqCst);
let failed = login(&issuer, /*callback_url*/ None).await?;
callback(
&failed.authorization_url(),
&issuer,
/*provider_error*/ false,
)
.await?;
let save_error = failed
.wait()
let save_error = complete_login(&issuer)
.await?
.commit_if(|| async { Some(()) })
.await

View File

@@ -127,17 +127,6 @@ async fn logout_invalidates_pending_login_across_processes() -> Result<()> {
Ok(())
}
async fn complete_login(issuer: &str) -> Result<EnterpriseOAuthCredentials> {
let handle = login(issuer, /*callback_url*/ None).await?;
callback(
&handle.authorization_url(),
issuer,
/*provider_error*/ false,
)
.await?;
handle.wait().await
}
fn stored(issuer: &str) -> Result<Option<StoredOAuthTokens>> {
crate::stored_oauth_credentials(
CREDENTIAL_NAME,