mirror of
https://github.com/openai/codex.git
synced 2026-09-08 15:50:34 +00:00
Cover OAuth refresh in local and remote MCP tests
Co-authored-by: Codex <noreply@openai.com>
This commit is contained in:
@@ -101,6 +101,7 @@ enum McpCallEvent {
|
||||
}
|
||||
|
||||
const REMOTE_MCP_ENVIRONMENT: &str = "remote";
|
||||
const LOCAL_MCP_ENVIRONMENT: &str = "local";
|
||||
|
||||
fn remote_aware_experimental_environment() -> Option<String> {
|
||||
// These tests run locally in normal CI and against the Docker-backed
|
||||
@@ -109,6 +110,10 @@ fn remote_aware_experimental_environment() -> Option<String> {
|
||||
std::env::var_os(remote_env_env_var()).map(|_| REMOTE_MCP_ENVIRONMENT.to_string())
|
||||
}
|
||||
|
||||
fn remote_only_experimental_environment() -> Option<String> {
|
||||
std::env::var_os(remote_env_env_var()).map(|_| REMOTE_MCP_ENVIRONMENT.to_string())
|
||||
}
|
||||
|
||||
/// Returns the stdio MCP test server command path for the active test placement.
|
||||
///
|
||||
/// Local test runs can execute the host-built test binary directly. Remote-aware
|
||||
@@ -1927,6 +1932,53 @@ struct StreamableHttpTestServer {
|
||||
process: StreamableHttpTestServerProcess,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum StreamableHttpTestServerBindMode {
|
||||
HostVisible,
|
||||
RemoteLoopbackOnly,
|
||||
}
|
||||
|
||||
struct StreamableHttpTestServerOptions<'a> {
|
||||
expected_env_value: &'a str,
|
||||
expected_bearer: Option<&'a str>,
|
||||
refreshed_access_token: Option<&'a str>,
|
||||
bind_mode: StreamableHttpTestServerBindMode,
|
||||
}
|
||||
|
||||
impl<'a> StreamableHttpTestServerOptions<'a> {
|
||||
fn host_visible(expected_env_value: &'a str, expected_bearer: Option<&'a str>) -> Self {
|
||||
Self {
|
||||
expected_env_value,
|
||||
expected_bearer,
|
||||
refreshed_access_token: None,
|
||||
bind_mode: StreamableHttpTestServerBindMode::HostVisible,
|
||||
}
|
||||
}
|
||||
|
||||
fn with_refreshed_access_token(mut self, refreshed_access_token: &'a str) -> Self {
|
||||
self.refreshed_access_token = Some(refreshed_access_token);
|
||||
self
|
||||
}
|
||||
|
||||
fn remote_loopback_only(mut self) -> Self {
|
||||
self.bind_mode = StreamableHttpTestServerBindMode::RemoteLoopbackOnly;
|
||||
self
|
||||
}
|
||||
|
||||
fn remote_loopback_only_if_remote(self, placement: OAuthRefreshPlacement) -> Self {
|
||||
match placement {
|
||||
OAuthRefreshPlacement::Local => self,
|
||||
OAuthRefreshPlacement::Remote => self.remote_loopback_only(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum OAuthTokenExpiry {
|
||||
Valid,
|
||||
Expired,
|
||||
}
|
||||
|
||||
/// Tracks whether the Streamable HTTP test server runs on the host or remotely.
|
||||
enum StreamableHttpTestServerProcess {
|
||||
Local(Child),
|
||||
@@ -2040,7 +2092,11 @@ async fn streamable_http_tool_call_round_trip() -> anyhow::Result<()> {
|
||||
// it is a host process.
|
||||
let expected_env_value = "propagated-env-http";
|
||||
let Some(http_server) =
|
||||
start_streamable_http_test_server(expected_env_value, /*expected_token*/ None).await?
|
||||
start_streamable_http_test_server(StreamableHttpTestServerOptions::host_visible(
|
||||
expected_env_value,
|
||||
/*expected_bearer*/ None,
|
||||
))
|
||||
.await?
|
||||
else {
|
||||
return Ok(());
|
||||
};
|
||||
@@ -2225,8 +2281,10 @@ async fn streamable_http_with_oauth_round_trip_impl() -> anyhow::Result<()> {
|
||||
let expected_token = "initial-access-token";
|
||||
let client_id = "test-client-id";
|
||||
let refresh_token = "initial-refresh-token";
|
||||
let Some(http_server) =
|
||||
start_streamable_http_test_server(expected_env_value, Some(expected_token)).await?
|
||||
let Some(http_server) = start_streamable_http_test_server(
|
||||
StreamableHttpTestServerOptions::host_visible(expected_env_value, Some(expected_token)),
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(());
|
||||
};
|
||||
@@ -2243,6 +2301,7 @@ async fn streamable_http_with_oauth_round_trip_impl() -> anyhow::Result<()> {
|
||||
client_id,
|
||||
expected_token,
|
||||
refresh_token,
|
||||
OAuthTokenExpiry::Valid,
|
||||
)?;
|
||||
|
||||
// Phase 4: configure Codex with the OAuth-backed Streamable HTTP MCP
|
||||
@@ -2361,10 +2420,235 @@ async fn streamable_http_with_oauth_round_trip_impl() -> anyhow::Result<()> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[serial(codex_home)]
|
||||
fn streamable_http_with_oauth_refresh_round_trip_local() -> anyhow::Result<()> {
|
||||
run_streamable_http_with_oauth_refresh_round_trip(OAuthRefreshPlacement::Local)
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[serial(codex_home)]
|
||||
fn streamable_http_with_oauth_refresh_round_trip_remote() -> anyhow::Result<()> {
|
||||
if remote_only_experimental_environment().is_none() {
|
||||
return Ok(());
|
||||
}
|
||||
run_streamable_http_with_oauth_refresh_round_trip(OAuthRefreshPlacement::Remote)
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum OAuthRefreshPlacement {
|
||||
Local,
|
||||
Remote,
|
||||
}
|
||||
|
||||
fn run_streamable_http_with_oauth_refresh_round_trip(
|
||||
placement: OAuthRefreshPlacement,
|
||||
) -> anyhow::Result<()> {
|
||||
const TEST_STACK_SIZE_BYTES: usize = 8 * 1024 * 1024;
|
||||
|
||||
let thread_name = match placement {
|
||||
OAuthRefreshPlacement::Local => "streamable_http_with_oauth_refresh_round_trip_local",
|
||||
OAuthRefreshPlacement::Remote => "streamable_http_with_oauth_refresh_round_trip_remote",
|
||||
};
|
||||
|
||||
let handle = std::thread::Builder::new()
|
||||
.name(thread_name.to_string())
|
||||
.stack_size(TEST_STACK_SIZE_BYTES)
|
||||
.spawn(move || -> anyhow::Result<()> {
|
||||
let runtime = tokio::runtime::Builder::new_multi_thread()
|
||||
.worker_threads(1)
|
||||
.enable_all()
|
||||
.build()?;
|
||||
runtime.block_on(streamable_http_with_oauth_refresh_round_trip_impl(
|
||||
placement,
|
||||
))
|
||||
})?;
|
||||
|
||||
match handle.join() {
|
||||
Ok(result) => result,
|
||||
Err(_) => Err(anyhow::anyhow!("{thread_name} thread panicked")),
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(clippy::expect_used)]
|
||||
async fn streamable_http_with_oauth_refresh_round_trip_impl(
|
||||
placement: OAuthRefreshPlacement,
|
||||
) -> anyhow::Result<()> {
|
||||
skip_if_no_network!(Ok(()));
|
||||
|
||||
let server = responses::start_mock_server().await;
|
||||
|
||||
let call_id = "call-790";
|
||||
let server_name = match placement {
|
||||
OAuthRefreshPlacement::Local => "rmcp_http_oauth_refresh_local",
|
||||
OAuthRefreshPlacement::Remote => "rmcp_http_oauth_refresh_remote",
|
||||
};
|
||||
let tool_name = format!("mcp__{server_name}__echo");
|
||||
let namespace = format!("mcp__{server_name}__");
|
||||
|
||||
mount_sse_once(
|
||||
&server,
|
||||
responses::sse(vec![
|
||||
responses::ev_response_created("resp-1"),
|
||||
responses::ev_function_call_with_namespace(
|
||||
call_id,
|
||||
&namespace,
|
||||
"echo",
|
||||
"{\"message\":\"ping\"}",
|
||||
),
|
||||
responses::ev_completed("resp-1"),
|
||||
]),
|
||||
)
|
||||
.await;
|
||||
mount_sse_once(
|
||||
&server,
|
||||
responses::sse(vec![
|
||||
responses::ev_assistant_message(
|
||||
"msg-1",
|
||||
"rmcp streamable http oauth refresh echo tool completed successfully.",
|
||||
),
|
||||
responses::ev_completed("resp-2"),
|
||||
]),
|
||||
)
|
||||
.await;
|
||||
|
||||
let expected_env_value = match placement {
|
||||
OAuthRefreshPlacement::Local => "propagated-env-http-oauth-refresh-local",
|
||||
OAuthRefreshPlacement::Remote => "propagated-env-http-oauth-refresh-remote",
|
||||
};
|
||||
let initial_access_token = "expired-access-token";
|
||||
let refreshed_access_token = match placement {
|
||||
OAuthRefreshPlacement::Local => "refreshed-access-token-local",
|
||||
OAuthRefreshPlacement::Remote => "refreshed-access-token-remote",
|
||||
};
|
||||
let refresh_token = "initial-refresh-token";
|
||||
// The remote case binds the test server to 127.0.0.1 inside the remote
|
||||
// container so the orchestrator cannot reach metadata or token endpoints
|
||||
// directly. If refresh still succeeds, it had to go through the selected
|
||||
// remote `HttpClient`.
|
||||
let Some(http_server) = start_streamable_http_test_server(
|
||||
StreamableHttpTestServerOptions::host_visible(
|
||||
expected_env_value,
|
||||
Some(refreshed_access_token),
|
||||
)
|
||||
.with_refreshed_access_token(refreshed_access_token)
|
||||
.remote_loopback_only_if_remote(placement),
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(());
|
||||
};
|
||||
let server_url = http_server.url().to_string();
|
||||
|
||||
let temp_home = Arc::new(tempdir()?);
|
||||
let _codex_home_guard = EnvVarGuard::set("CODEX_HOME", temp_home.path().as_os_str());
|
||||
write_fallback_oauth_tokens(
|
||||
temp_home.path(),
|
||||
server_name,
|
||||
&server_url,
|
||||
"test-client-id",
|
||||
initial_access_token,
|
||||
refresh_token,
|
||||
OAuthTokenExpiry::Expired,
|
||||
)?;
|
||||
|
||||
let experimental_environment = match placement {
|
||||
OAuthRefreshPlacement::Local => Some(LOCAL_MCP_ENVIRONMENT.to_string()),
|
||||
OAuthRefreshPlacement::Remote => Some(REMOTE_MCP_ENVIRONMENT.to_string()),
|
||||
};
|
||||
let fixture = test_codex()
|
||||
.with_home(temp_home.clone())
|
||||
.with_config(move |config| {
|
||||
config.mcp_oauth_credentials_store_mode = serde_json::from_value(json!("file"))
|
||||
.expect("`file` should deserialize as OAuthCredentialsStoreMode");
|
||||
insert_mcp_server(
|
||||
config,
|
||||
server_name,
|
||||
McpServerTransportConfig::StreamableHttp {
|
||||
url: server_url,
|
||||
bearer_token_env_var: None,
|
||||
http_headers: None,
|
||||
env_http_headers: None,
|
||||
},
|
||||
TestMcpServerOptions {
|
||||
experimental_environment,
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
})
|
||||
.build_remote_aware(&server)
|
||||
.await?;
|
||||
let session_model = fixture.session_configured.model.clone();
|
||||
|
||||
wait_for_mcp_tool(&fixture, &tool_name).await?;
|
||||
|
||||
fixture
|
||||
.codex
|
||||
.submit(Op::UserTurn {
|
||||
items: vec![UserInput::Text {
|
||||
text: "call the rmcp streamable http oauth echo tool".into(),
|
||||
text_elements: Vec::new(),
|
||||
}],
|
||||
final_output_json_schema: None,
|
||||
cwd: fixture.cwd.path().to_path_buf(),
|
||||
approval_policy: AskForApproval::Never,
|
||||
approvals_reviewer: None,
|
||||
sandbox_policy: SandboxPolicy::new_read_only_policy(),
|
||||
model: session_model,
|
||||
effort: None,
|
||||
summary: None,
|
||||
service_tier: None,
|
||||
collaboration_mode: None,
|
||||
personality: None,
|
||||
environments: None,
|
||||
})
|
||||
.await?;
|
||||
|
||||
let end_event = wait_for_event(&fixture.codex, |ev| {
|
||||
matches!(ev, EventMsg::McpToolCallEnd(_))
|
||||
})
|
||||
.await;
|
||||
let EventMsg::McpToolCallEnd(end) = end_event else {
|
||||
unreachable!("event guard guarantees McpToolCallEnd");
|
||||
};
|
||||
|
||||
let result = end
|
||||
.result
|
||||
.as_ref()
|
||||
.expect("rmcp echo tool should return success");
|
||||
assert_eq!(result.is_error, Some(false));
|
||||
let structured = result
|
||||
.structured_content
|
||||
.as_ref()
|
||||
.expect("structured content");
|
||||
let Value::Object(map) = structured else {
|
||||
panic!("structured content should be an object: {structured:?}");
|
||||
};
|
||||
assert_eq!(
|
||||
map.get("echo").and_then(Value::as_str),
|
||||
Some("ECHOING: ping")
|
||||
);
|
||||
assert_eq!(
|
||||
map.get("env").and_then(Value::as_str),
|
||||
Some(expected_env_value)
|
||||
);
|
||||
|
||||
wait_for_event(&fixture.codex, |ev| matches!(ev, EventMsg::TurnComplete(_))).await;
|
||||
server.verify().await;
|
||||
|
||||
assert_eq!(
|
||||
read_fallback_oauth_access_token(temp_home.path())?,
|
||||
refreshed_access_token
|
||||
);
|
||||
|
||||
http_server.shutdown().await;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Starts the Streamable HTTP MCP test server in the active test placement.
|
||||
async fn start_streamable_http_test_server(
|
||||
expected_env_value: &str,
|
||||
expected_token: Option<&str>,
|
||||
options: StreamableHttpTestServerOptions<'_>,
|
||||
) -> anyhow::Result<Option<StreamableHttpTestServer>> {
|
||||
let rmcp_http_server_bin = match cargo_bin("test_streamable_http_server") {
|
||||
Ok(path) => path,
|
||||
@@ -2379,8 +2663,7 @@ async fn start_streamable_http_test_server(
|
||||
start_remote_streamable_http_test_server(
|
||||
&container_name,
|
||||
&rmcp_http_server_bin,
|
||||
expected_env_value,
|
||||
expected_token,
|
||||
options,
|
||||
)
|
||||
.await?,
|
||||
));
|
||||
@@ -2396,10 +2679,13 @@ async fn start_streamable_http_test_server(
|
||||
command
|
||||
.kill_on_drop(true)
|
||||
.env("MCP_STREAMABLE_HTTP_BIND_ADDR", &bind_addr)
|
||||
.env("MCP_TEST_VALUE", expected_env_value);
|
||||
if let Some(expected_token) = expected_token {
|
||||
.env("MCP_TEST_VALUE", options.expected_env_value);
|
||||
if let Some(expected_token) = options.expected_bearer {
|
||||
command.env("MCP_EXPECT_BEARER", expected_token);
|
||||
}
|
||||
if let Some(refreshed_access_token) = options.refreshed_access_token {
|
||||
command.env("MCP_REFRESHED_ACCESS_TOKEN", refreshed_access_token);
|
||||
}
|
||||
let mut child = command.spawn()?;
|
||||
|
||||
wait_for_local_streamable_http_server(&mut child, &server_url, Duration::from_secs(5)).await?;
|
||||
@@ -2413,8 +2699,7 @@ async fn start_streamable_http_test_server(
|
||||
async fn start_remote_streamable_http_test_server(
|
||||
container_name: &str,
|
||||
rmcp_http_server_bin: &Path,
|
||||
expected_env_value: &str,
|
||||
expected_token: Option<&str>,
|
||||
options: StreamableHttpTestServerOptions<'_>,
|
||||
) -> anyhow::Result<StreamableHttpTestServer> {
|
||||
let remote_path = copy_binary_to_remote_env(
|
||||
container_name,
|
||||
@@ -2426,20 +2711,32 @@ async fn start_remote_streamable_http_test_server(
|
||||
let mut env_assignments = vec![
|
||||
format!(
|
||||
"MCP_STREAMABLE_HTTP_BIND_ADDR={}",
|
||||
sh_single_quote("0.0.0.0:0")
|
||||
sh_single_quote(match options.bind_mode {
|
||||
StreamableHttpTestServerBindMode::HostVisible => "0.0.0.0:0",
|
||||
StreamableHttpTestServerBindMode::RemoteLoopbackOnly => "127.0.0.1:0",
|
||||
})
|
||||
),
|
||||
format!(
|
||||
"MCP_STREAMABLE_HTTP_BOUND_ADDR_FILE={}",
|
||||
sh_single_quote(&bound_addr_file)
|
||||
),
|
||||
format!("MCP_TEST_VALUE={}", sh_single_quote(expected_env_value)),
|
||||
format!(
|
||||
"MCP_TEST_VALUE={}",
|
||||
sh_single_quote(options.expected_env_value)
|
||||
),
|
||||
];
|
||||
if let Some(expected_token) = expected_token {
|
||||
if let Some(expected_token) = options.expected_bearer {
|
||||
env_assignments.push(format!(
|
||||
"MCP_EXPECT_BEARER={}",
|
||||
sh_single_quote(expected_token)
|
||||
));
|
||||
}
|
||||
if let Some(refreshed_access_token) = options.refreshed_access_token {
|
||||
env_assignments.push(format!(
|
||||
"MCP_REFRESHED_ACCESS_TOKEN={}",
|
||||
sh_single_quote(refreshed_access_token)
|
||||
));
|
||||
}
|
||||
|
||||
let script = format!(
|
||||
"{} nohup {} > {} 2>&1 < /dev/null & echo $!",
|
||||
@@ -2469,15 +2766,15 @@ async fn start_remote_streamable_http_test_server(
|
||||
let remote_bind_addr =
|
||||
wait_for_remote_bound_addr(container_name, &bound_addr_file, Duration::from_secs(5))
|
||||
.await?;
|
||||
let container_ip = remote_container_ip(container_name)?;
|
||||
let server_url = format!("http://{}:{}/mcp", container_ip, remote_bind_addr.port());
|
||||
let server_host = match options.bind_mode {
|
||||
StreamableHttpTestServerBindMode::HostVisible => remote_container_ip(container_name)?,
|
||||
StreamableHttpTestServerBindMode::RemoteLoopbackOnly => "127.0.0.1".to_string(),
|
||||
};
|
||||
let server_url = format!("http://{}:{}/mcp", server_host, remote_bind_addr.port());
|
||||
// The orchestrator can see the Docker container IP, but the behavior under
|
||||
// test is whether the remote-side MCP client can reach it. Probe through
|
||||
// remote HTTP before handing the URL to the Codex fixture.
|
||||
wait_for_remote_streamable_http_server(&server_url, Duration::from_secs(5)).await?;
|
||||
if expected_token.is_some() {
|
||||
wait_for_streamable_http_metadata(&server_url, Duration::from_secs(5)).await?;
|
||||
}
|
||||
|
||||
Ok(StreamableHttpTestServer {
|
||||
server_url,
|
||||
@@ -2660,50 +2957,6 @@ async fn wait_for_remote_streamable_http_server(
|
||||
}
|
||||
}
|
||||
|
||||
/// Waits for OAuth metadata from the host-side test process.
|
||||
async fn wait_for_streamable_http_metadata(
|
||||
server_url: &str,
|
||||
timeout: Duration,
|
||||
) -> anyhow::Result<()> {
|
||||
let deadline = Instant::now() + timeout;
|
||||
let metadata_url = streamable_http_metadata_url(server_url);
|
||||
let client = Client::builder().no_proxy().build()?;
|
||||
loop {
|
||||
let remaining = deadline.saturating_duration_since(Instant::now());
|
||||
if remaining.is_zero() {
|
||||
return Err(anyhow::anyhow!(
|
||||
"timed out waiting for streamable HTTP server metadata at {metadata_url}: deadline reached"
|
||||
));
|
||||
}
|
||||
|
||||
match tokio::time::timeout(remaining, client.get(&metadata_url).send()).await {
|
||||
Ok(Ok(response)) if response.status() == StatusCode::OK => return Ok(()),
|
||||
Ok(Ok(response)) => {
|
||||
if Instant::now() >= deadline {
|
||||
return Err(anyhow::anyhow!(
|
||||
"timed out waiting for streamable HTTP server metadata at {metadata_url}: HTTP {}",
|
||||
response.status()
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(Err(error)) => {
|
||||
if Instant::now() >= deadline {
|
||||
return Err(anyhow::anyhow!(
|
||||
"timed out waiting for streamable HTTP server metadata at {metadata_url}: {error}"
|
||||
));
|
||||
}
|
||||
}
|
||||
Err(_) => {
|
||||
return Err(anyhow::anyhow!(
|
||||
"timed out waiting for streamable HTTP server metadata at {metadata_url}: request timed out"
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
sleep(Duration::from_millis(50)).await;
|
||||
}
|
||||
}
|
||||
|
||||
/// Builds the OAuth metadata URL for the test Streamable HTTP MCP endpoint.
|
||||
fn streamable_http_metadata_url(server_url: &str) -> String {
|
||||
let base_url = server_url.strip_suffix("/mcp").unwrap_or(server_url);
|
||||
@@ -2717,12 +2970,18 @@ fn write_fallback_oauth_tokens(
|
||||
client_id: &str,
|
||||
access_token: &str,
|
||||
refresh_token: &str,
|
||||
expiry: OAuthTokenExpiry,
|
||||
) -> anyhow::Result<()> {
|
||||
let expires_at = SystemTime::now()
|
||||
.checked_add(Duration::from_secs(3600))
|
||||
.ok_or_else(|| anyhow::anyhow!("failed to compute expiry time"))?
|
||||
.duration_since(UNIX_EPOCH)?
|
||||
.as_millis() as u64;
|
||||
let expires_at = match expiry {
|
||||
OAuthTokenExpiry::Valid => SystemTime::now()
|
||||
.checked_add(Duration::from_secs(3600))
|
||||
.ok_or_else(|| anyhow::anyhow!("failed to compute expiry time"))?,
|
||||
OAuthTokenExpiry::Expired => SystemTime::now()
|
||||
.checked_sub(Duration::from_secs(3600))
|
||||
.ok_or_else(|| anyhow::anyhow!("failed to compute expiry time"))?,
|
||||
}
|
||||
.duration_since(UNIX_EPOCH)?
|
||||
.as_millis() as u64;
|
||||
|
||||
let store = serde_json::json!({
|
||||
"stub": {
|
||||
@@ -2741,6 +3000,17 @@ fn write_fallback_oauth_tokens(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn read_fallback_oauth_access_token(home: &Path) -> anyhow::Result<String> {
|
||||
let file_path = home.join(".credentials.json");
|
||||
let store: Value = serde_json::from_slice(&fs::read(&file_path)?)?;
|
||||
store
|
||||
.get("stub")
|
||||
.and_then(|stub| stub.get("access_token"))
|
||||
.and_then(Value::as_str)
|
||||
.map(ToOwned::to_owned)
|
||||
.ok_or_else(|| anyhow::anyhow!("missing fallback OAuth access token"))
|
||||
}
|
||||
|
||||
struct EnvVarGuard {
|
||||
key: &'static str,
|
||||
original: Option<OsString>,
|
||||
|
||||
@@ -62,6 +62,7 @@ const MEMO_URI: &str = "memo://codex/example-note";
|
||||
const MEMO_CONTENT: &str = "This is a sample MCP resource served by the rmcp test server.";
|
||||
const MCP_SESSION_ID_HEADER: &str = "mcp-session-id";
|
||||
const SESSION_POST_FAILURE_CONTROL_PATH: &str = "/test/control/session-post-failure";
|
||||
const OAUTH_TOKEN_PATH: &str = "/oauth/token";
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
struct SessionFailureState {
|
||||
@@ -116,6 +117,8 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
fs::write(bound_addr_file, actual_bind_addr.to_string())?;
|
||||
}
|
||||
eprintln!("starting rmcp streamable http test server on http://{actual_bind_addr}/mcp");
|
||||
let refreshed_access_token = std::env::var("MCP_REFRESHED_ACCESS_TOKEN")
|
||||
.unwrap_or_else(|_| "refreshed-access-token".to_string());
|
||||
|
||||
let router = Router::new()
|
||||
.route(
|
||||
@@ -146,6 +149,28 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
}
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
OAUTH_TOKEN_PATH,
|
||||
post({
|
||||
move || async move {
|
||||
let refreshed_access_token = refreshed_access_token.clone();
|
||||
#[expect(clippy::expect_used)]
|
||||
Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(CONTENT_TYPE, "application/json")
|
||||
.body(Body::from(
|
||||
serde_json::to_vec(&json!({
|
||||
"access_token": refreshed_access_token,
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600_u64,
|
||||
"refresh_token": "refreshed-refresh-token",
|
||||
}))
|
||||
.expect("failed to serialize token response"),
|
||||
))
|
||||
.expect("valid token response")
|
||||
}
|
||||
}),
|
||||
)
|
||||
.nest_service(
|
||||
"/mcp",
|
||||
StreamableHttpService::new(
|
||||
@@ -386,7 +411,8 @@ async fn require_bearer(
|
||||
request: Request<Body>,
|
||||
next: Next,
|
||||
) -> Result<Response, StatusCode> {
|
||||
if request.uri().path().contains("/.well-known/") {
|
||||
if request.uri().path().contains("/.well-known/") || request.uri().path().starts_with("/oauth/")
|
||||
{
|
||||
return Ok(next.run(request).await);
|
||||
}
|
||||
if request
|
||||
|
||||
@@ -933,6 +933,10 @@ async fn create_oauth_transport_and_runtime(
|
||||
StreamableHttpClientTransport<AuthClient<StreamableHttpClientAdapter>>,
|
||||
OAuthPersistor,
|
||||
)> {
|
||||
// `create_oauth_http_setup` returns both the initialized RMCP
|
||||
// `AuthorizationManager` and the proxy bridge that keeps later token
|
||||
// refresh requests on the selected `HttpClient`. Keep both pieces together
|
||||
// here so the transport and persistor share the same OAuth runtime.
|
||||
let oauth_http_setup = create_oauth_http_setup(
|
||||
url,
|
||||
&default_headers,
|
||||
|
||||
Reference in New Issue
Block a user