Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
142 changes: 131 additions & 11 deletions codex-rs/app-server/tests/suite/v2/mcp_resource.rs
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::time::Duration;

use anyhow::Result;
Expand Down Expand Up @@ -79,11 +81,13 @@ const SKILL_REFERENCE_CONTENTS: &str =
"# Deploy reference\n\nUse the orchestrator deployment API.\n";
const SKILLS_LIST_CALL_ID: &str = "skills-list";
const SKILLS_READ_CALL_ID: &str = "skills-read";
const SKILLS_READ_AGAIN_CALL_ID: &str = "skills-read-again";

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn mcp_resource_read_returns_resource_contents() -> Result<()> {
let responses_server = responses::start_mock_server().await;
let (apps_server_url, apps_server_handle) = start_resource_apps_mcp_server().await?;
let (apps_server_url, _apps_server_calls, apps_server_handle) =
start_resource_apps_mcp_server().await?;
let responses_server_uri = responses_server.uri();
let (_codex_home, mut mcp) =
start_resource_test_app_server(&apps_server_url, &responses_server_uri).await?;
Expand Down Expand Up @@ -126,7 +130,8 @@ async fn mcp_resource_read_returns_resource_contents() -> Result<()> {
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn orchestrator_skill_can_read_referenced_resource_without_an_executor() -> Result<()> {
let responses_server = responses::start_mock_server().await;
let (apps_server_url, apps_server_handle) = start_resource_apps_mcp_server().await?;
let (apps_server_url, apps_server_calls, apps_server_handle) =
start_resource_apps_mcp_server().await?;
let responses_server_uri = responses_server.uri();
let (_codex_home, mut mcp) =
start_resource_test_app_server(&apps_server_url, &responses_server_uri).await?;
Expand Down Expand Up @@ -180,17 +185,39 @@ async fn orchestrator_skill_can_read_referenced_resource_without_an_executor() -
),
responses::ev_completed("resp-skills-read"),
]),
responses::sse(vec![
responses::ev_response_created("resp-skills-read-again"),
responses::ev_function_call_with_namespace(
SKILLS_READ_AGAIN_CALL_ID,
"skills",
"read",
&json!({
"authority": {
"kind": "orchestrator",
},
"package": SKILL_RESOURCE_URI,
"resource": SKILL_REFERENCE_URI,
})
.to_string(),
),
responses::ev_completed("resp-skills-read-again"),
]),
responses::sse(vec![
responses::ev_response_created("resp-orchestrator-skill"),
responses::ev_assistant_message("msg-orchestrator-skill", "Done"),
responses::ev_completed("resp-orchestrator-skill"),
]),
responses::sse(vec![
responses::ev_response_created("resp-orchestrator-skill-after-refresh"),
responses::ev_assistant_message("msg-orchestrator-skill-after-refresh", "Done"),
responses::ev_completed("resp-orchestrator-skill-after-refresh"),
]),
],
)
.await;
let turn_start_id = mcp
.send_turn_start_request(TurnStartParams {
thread_id: thread.id,
thread_id: thread.id.clone(),
input: vec![UserInput::Text {
text: format!("Use ${SKILL_NAME}"),
text_elements: Vec::new(),
Expand All @@ -210,7 +237,7 @@ async fn orchestrator_skill_can_read_referenced_resource_without_an_executor() -
.await??;

let requests = response_mock.requests();
assert_eq!(requests.len(), 3);
assert_eq!(requests.len(), 4);
let first_request = &requests[0];
assert!(first_request.tool_by_name("skills", "list").is_some());
assert!(first_request.tool_by_name("skills", "read").is_some());
Expand Down Expand Up @@ -276,6 +303,61 @@ async fn orchestrator_skill_can_read_referenced_resource_without_an_executor() -
"contents": SKILL_REFERENCE_CONTENTS,
})
);
let repeated_read_output = requests[3]
.function_call_output_text(SKILLS_READ_AGAIN_CALL_ID)
.ok_or_else(|| {
anyhow::anyhow!("repeated skills.read output should be sent to the model")
})?;
assert_eq!(read_output, repeated_read_output);
assert_eq!(
ResourceAppsMcpCallCounts {
list_resources: 3,
main_prompt_reads: 1,
reference_reads: 1,
},
apps_server_calls.snapshot()
);

let refresh_request_id = mcp
.send_raw_request("config/mcpServer/reload", /*params*/ None)
.await?;
timeout(
DEFAULT_READ_TIMEOUT,
mcp.read_stream_until_response_message(RequestId::Integer(refresh_request_id)),
)
.await??;

let refreshed_turn_start_id = mcp
.send_turn_start_request(TurnStartParams {
thread_id: thread.id,
input: vec![UserInput::Text {
text: format!("Use ${SKILL_NAME} after refreshing MCP"),
text_elements: Vec::new(),
}],
..Default::default()
})
.await?;
timeout(
DEFAULT_READ_TIMEOUT,
mcp.read_stream_until_response_message(RequestId::Integer(refreshed_turn_start_id)),
)
.await??;
timeout(
DEFAULT_READ_TIMEOUT,
mcp.read_stream_until_notification_message("turn/completed"),
)
.await??;

let requests = response_mock.requests();
assert_eq!(requests.len(), 5);
assert_eq!(
ResourceAppsMcpCallCounts {
list_resources: 6,
main_prompt_reads: 2,
reference_reads: 1,
},
apps_server_calls.snapshot()
);
apps_server_handle.abort();
let _ = apps_server_handle.await;
Ok(())
Expand All @@ -284,7 +366,8 @@ async fn orchestrator_skill_can_read_referenced_resource_without_an_executor() -
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn local_executor_does_not_expose_orchestrator_skills() -> Result<()> {
let responses_server = responses::start_mock_server().await;
let (apps_server_url, apps_server_handle) = start_resource_apps_mcp_server().await?;
let (apps_server_url, _apps_server_calls, apps_server_handle) =
start_resource_apps_mcp_server().await?;
let responses_server_uri = responses_server.uri();
let (_codex_home, mut mcp) =
start_resource_test_app_server(&apps_server_url, &responses_server_uri).await?;
Expand Down Expand Up @@ -355,7 +438,8 @@ async fn local_executor_does_not_expose_orchestrator_skills() -> Result<()> {

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn mcp_resource_read_returns_resource_contents_without_thread() -> Result<()> {
let (apps_server_url, apps_server_handle) = start_resource_apps_mcp_server().await?;
let (apps_server_url, _apps_server_calls, apps_server_handle) =
start_resource_apps_mcp_server().await?;

let codex_home = TempDir::new()?;
std::fs::write(
Expand Down Expand Up @@ -514,13 +598,20 @@ stream_max_retries = 0
Ok((codex_home, mcp))
}

async fn start_resource_apps_mcp_server() -> Result<(String, JoinHandle<()>)> {
async fn start_resource_apps_mcp_server()
-> Result<(String, Arc<ResourceAppsMcpCalls>, JoinHandle<()>)> {
let listener = TcpListener::bind("127.0.0.1:0").await?;
let addr = listener.local_addr()?;
let apps_server_url = format!("http://{addr}");
let calls = Arc::new(ResourceAppsMcpCalls::default());
let server_calls = Arc::clone(&calls);

let mcp_service = StreamableHttpService::new(
move || Ok(ResourceAppsMcpServer),
move || {
Ok(ResourceAppsMcpServer {
calls: Arc::clone(&server_calls),
})
},
Arc::new(LocalSessionManager::default()),
StreamableHttpServerConfig::default(),
);
Expand All @@ -529,7 +620,7 @@ async fn start_resource_apps_mcp_server() -> Result<(String, JoinHandle<()>)> {
let _ = axum::serve(listener, router).await;
});

Ok((apps_server_url, apps_server_handle))
Ok((apps_server_url, calls, apps_server_handle))
}

fn expected_resource_read_response() -> McpResourceReadResponse {
Expand All @@ -551,8 +642,34 @@ fn expected_resource_read_response() -> McpResourceReadResponse {
}
}

#[derive(Clone, Default)]
struct ResourceAppsMcpServer;
#[derive(Debug, Default)]
struct ResourceAppsMcpCalls {
list_resources: AtomicUsize,
main_prompt_reads: AtomicUsize,
reference_reads: AtomicUsize,
}

impl ResourceAppsMcpCalls {
fn snapshot(&self) -> ResourceAppsMcpCallCounts {
ResourceAppsMcpCallCounts {
list_resources: self.list_resources.load(Ordering::Relaxed),
main_prompt_reads: self.main_prompt_reads.load(Ordering::Relaxed),
reference_reads: self.reference_reads.load(Ordering::Relaxed),
}
}
}

#[derive(Debug, PartialEq, Eq)]
struct ResourceAppsMcpCallCounts {
list_resources: usize,
main_prompt_reads: usize,
reference_reads: usize,
}

#[derive(Clone)]
struct ResourceAppsMcpServer {
calls: Arc<ResourceAppsMcpCalls>,
}

impl ServerHandler for ResourceAppsMcpServer {
fn get_info(&self) -> ServerInfo {
Expand All @@ -565,6 +682,7 @@ impl ServerHandler for ResourceAppsMcpServer {
request: Option<PaginatedRequestParams>,
_context: RequestContext<RoleServer>,
) -> Result<ListResourcesResult, rmcp::ErrorData> {
self.calls.list_resources.fetch_add(1, Ordering::Relaxed);
let cursor = request.and_then(|request| request.cursor);
if cursor.is_none() {
return Ok(ListResourcesResult {
Expand Down Expand Up @@ -614,6 +732,7 @@ impl ServerHandler for ResourceAppsMcpServer {
) -> Result<ReadResourceResult, rmcp::ErrorData> {
let uri = request.uri;
if uri == SKILL_MAIN_PROMPT_URI {
self.calls.main_prompt_reads.fetch_add(1, Ordering::Relaxed);
return Ok(ReadResourceResult::new(vec![
ResourceContents::TextResourceContents {
uri: SKILL_MAIN_PROMPT_URI.to_string(),
Expand All @@ -624,6 +743,7 @@ impl ServerHandler for ResourceAppsMcpServer {
]));
}
if uri == SKILL_REFERENCE_URI {
self.calls.reference_reads.fetch_add(1, Ordering::Relaxed);
return Ok(ReadResourceResult::new(vec![
ResourceContents::TextResourceContents {
uri: SKILL_REFERENCE_URI.to_string(),
Expand Down
1 change: 1 addition & 0 deletions codex-rs/codex-mcp/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ pub use elicitation::ElicitationReviewRequest;
pub use elicitation::ElicitationReviewer;
pub use elicitation::ElicitationReviewerHandle;
pub use resource_client::McpResourceClient;
pub use resource_client::McpResourceClientCacheKey;
pub use resource_client::McpResourcePage;
pub use resource_client::McpResourceReadResult;
pub use rmcp_client::MCP_SANDBOX_STATE_META_CAPABILITY;
Expand Down
18 changes: 18 additions & 0 deletions codex-rs/codex-mcp/src/resource_client.rs
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
use std::sync::Arc;
use std::sync::Weak;

use anyhow::Context;
use anyhow::Result;
Expand Down Expand Up @@ -35,6 +36,18 @@ pub struct McpResourceClient {
manager: Arc<ArcSwap<McpConnectionManager>>,
}

/// Opaque identity for the manager currently used by an MCP resource client.
#[derive(Clone)]
pub struct McpResourceClientCacheKey(Weak<McpConnectionManager>);

impl PartialEq for McpResourceClientCacheKey {
fn eq(&self, other: &Self) -> bool {
self.0.ptr_eq(&other.0)
}
}

impl Eq for McpResourceClientCacheKey {}

impl std::fmt::Debug for McpResourceClient {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
Expand All @@ -49,6 +62,11 @@ impl McpResourceClient {
Self { manager }
}

/// Returns an identity that changes whenever the published manager changes.
pub fn cache_key(&self) -> McpResourceClientCacheKey {
McpResourceClientCacheKey(Arc::downgrade(&self.manager.load_full()))
}

/// Returns whether the current manager contains the named server.
///
/// This does not wait for server startup or imply that startup succeeded.
Expand Down
37 changes: 25 additions & 12 deletions codex-rs/ext/skills/src/extension.rs
Original file line number Diff line number Diff line change
Expand Up @@ -150,17 +150,19 @@ where
session_store: &ExtensionData,
thread_store: &ExtensionData,
) -> Vec<Arc<dyn ToolExecutor<ToolCall>>> {
let Some(thread_state) = thread_store.get::<SkillsThreadState>() else {
return Vec::new();
};
if !self.providers.has_orchestrator_provider()
|| !thread_store
.get::<SkillsThreadState>()
.is_some_and(|state| state.orchestrator_skills_enabled())
|| !thread_state.orchestrator_skills_enabled()
{
return Vec::new();
}

skill_tools(
self.providers.clone(),
session_store.get::<McpResourceClient>(),
thread_state,
)
}
}
Expand Down Expand Up @@ -215,7 +217,12 @@ where
let mut injected_host_skill_prompts = InjectedHostSkillPrompts::default();
for entry in &selected_entries {
match self
.read_main_prompt(entry, host_loaded_skills.clone(), session_store)
.read_main_prompt(
entry,
host_loaded_skills.clone(),
session_store,
&thread_state,
)
.await
{
Ok(read_result) => {
Expand Down Expand Up @@ -292,12 +299,14 @@ impl<C> SkillsExtension<C> {
) -> SkillCatalog {
let include_orchestrator_skills = query.include_orchestrator_skills;
let orchestrator_query = query.clone();
let mcp_resources = orchestrator_query.mcp_resources.clone();
query.include_orchestrator_skills = false;

let mut catalog = self.providers.list_for_turn(query).await;
if include_orchestrator_skills {
let orchestrator_catalog = thread_state
.orchestrator_catalog_snapshot(
mcp_resources.as_deref(),
self.providers
.list_orchestrator_for_turn(orchestrator_query),
)
Expand All @@ -312,15 +321,19 @@ impl<C> SkillsExtension<C> {
entry: &SkillCatalogEntry,
host_loaded_skills: Option<Arc<HostLoadedSkills>>,
session_store: &ExtensionData,
thread_state: &SkillsThreadState,
) -> Result<SkillReadResult, String> {
self.providers
.read(SkillReadRequest {
authority: entry.authority.clone(),
package: entry.id.clone(),
resource: entry.main_prompt.clone(),
host: host_loaded_skills,
mcp_resources: session_store.get::<McpResourceClient>(),
})
thread_state
.read_skill(
&self.providers,
SkillReadRequest {
authority: entry.authority.clone(),
package: entry.id.clone(),
resource: entry.main_prompt.clone(),
host: host_loaded_skills,
mcp_resources: session_store.get::<McpResourceClient>(),
},
)
.await
.map_err(|err| err.message)
}
Expand Down
Loading
Loading