Skip to content
115 changes: 115 additions & 0 deletions rust/src/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1944,6 +1944,9 @@ pub struct SessionConfig {
pub session_id: Option<SessionId>,
/// Model to use (e.g. `"gpt-4"`, `"claude-sonnet-4"`).
pub model: Option<String>,
/// Exact model identifiers permitted for this session. When unset, the SDK
/// does not restrict model selection.
pub allowed_models: Option<Vec<String>>,
/// Application name sent as `User-Agent` context.
pub client_name: Option<String>,
/// Reasoning effort level (e.g. `"low"`, `"medium"`, `"high"`).
Expand Down Expand Up @@ -2302,6 +2305,7 @@ impl std::fmt::Debug for SessionConfig {
f.debug_struct("SessionConfig")
.field("session_id", &self.session_id)
.field("model", &self.model)
.field("allowed_models", &self.allowed_models)
.field("client_name", &self.client_name)
.field("reasoning_effort", &self.reasoning_effort)
.field("reasoning_summary", &self.reasoning_summary)
Expand Down Expand Up @@ -2449,6 +2453,7 @@ impl Default for SessionConfig {
Self {
session_id: None,
model: None,
allowed_models: None,
client_name: None,
reasoning_effort: None,
reasoning_summary: None,
Expand Down Expand Up @@ -2620,6 +2625,7 @@ impl SessionConfig {
let wire = crate::wire::SessionCreateWire {
session_id,
model: self.model,
allowed_models: self.allowed_models,
client_name: self.client_name,
reasoning_effort: self.reasoning_effort,
reasoning_summary: self.reasoning_summary,
Expand Down Expand Up @@ -2844,6 +2850,16 @@ impl SessionConfig {
self
}

/// Set the exact model identifiers permitted for this session.
pub fn with_allowed_models<I, S>(mut self, allowed_models: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
self.allowed_models = Some(allowed_models.into_iter().map(Into::into).collect());
self
}

/// Set the application name sent as `User-Agent` context.
pub fn with_client_name(mut self, name: impl Into<String>) -> Self {
self.client_name = Some(name.into());
Expand Down Expand Up @@ -3402,6 +3418,9 @@ pub struct ResumeSessionConfig {
/// Model to use for this session (e.g. `"gpt-4"`, `"claude-sonnet-4"`).
/// Can change the model when resuming.
pub model: Option<String>,
/// Exact model identifiers permitted for the resumed session. When unset,
/// the SDK does not restrict model selection.
pub allowed_models: Option<Vec<String>>,
/// Application name sent as User-Agent context.
pub client_name: Option<String>,
/// Desired reasoning effort to apply after resuming the session.
Expand Down Expand Up @@ -3672,6 +3691,7 @@ impl std::fmt::Debug for ResumeSessionConfig {
f.debug_struct("ResumeSessionConfig")
.field("session_id", &self.session_id)
.field("model", &self.model)
.field("allowed_models", &self.allowed_models)
.field("client_name", &self.client_name)
.field("reasoning_effort", &self.reasoning_effort)
.field("reasoning_summary", &self.reasoning_summary)
Expand Down Expand Up @@ -3862,6 +3882,7 @@ impl ResumeSessionConfig {
let wire = crate::wire::SessionResumeWire {
session_id: self.session_id,
model: self.model,
allowed_models: self.allowed_models,
client_name: self.client_name,
reasoning_effort: self.reasoning_effort,
reasoning_summary: self.reasoning_summary,
Expand Down Expand Up @@ -3971,6 +3992,7 @@ impl ResumeSessionConfig {
Self {
session_id,
model: None,
allowed_models: None,
client_name: None,
reasoning_effort: None,
reasoning_summary: None,
Expand Down Expand Up @@ -4166,6 +4188,16 @@ impl ResumeSessionConfig {
self
}

/// Set the exact model identifiers permitted for the resumed session.
pub fn with_allowed_models<I, S>(mut self, allowed_models: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
self.allowed_models = Some(allowed_models.into_iter().map(Into::into).collect());
self
}

/// Set the application name sent as `User-Agent` context.
pub fn with_client_name(mut self, name: impl Into<String>) -> Self {
self.client_name = Some(name.into());
Expand Down Expand Up @@ -6528,6 +6560,89 @@ mod tests {
assert!(json.get("askUserVariant").is_none());
}

#[test]
fn session_config_allowed_models_builder_debug_and_wire() {
let default = SessionConfig::default();
assert_eq!(default.allowed_models, None);
assert!(format!("{default:?}").contains("allowed_models: None"));

let config = SessionConfig::default().with_allowed_models(["gpt-5", "claude-sonnet-5"]);
assert_eq!(
config.allowed_models.as_deref(),
Some(&["gpt-5".to_string(), "claude-sonnet-5".to_string()][..])
);
assert!(
format!("{config:?}")
.contains("allowed_models: Some([\"gpt-5\", \"claude-sonnet-5\"])")
);

let (wire, _) = config
.into_wire(Some(SessionId::from("allowed-models-create")))
.expect("allowed models do not add SDK validation");
let json = serde_json::to_value(&wire).unwrap();
assert_eq!(json["allowedModels"], json!(["gpt-5", "claude-sonnet-5"]));

let (default_wire, _) = SessionConfig::default()
.into_wire(Some(SessionId::from("unrestricted-create")))
.expect("default config has no duplicate handlers");
let default_json = serde_json::to_value(&default_wire).unwrap();
assert!(default_json.get("allowedModels").is_none());
}

#[test]
fn resume_session_config_allowed_models_builder_debug_and_wire() {
let default = ResumeSessionConfig::new(SessionId::from("unrestricted-resume"));
assert_eq!(default.allowed_models, None);
assert!(format!("{default:?}").contains("allowed_models: None"));

let config = ResumeSessionConfig::new(SessionId::from("allowed-models-resume"))
.with_allowed_models(vec!["gpt-5".to_string(), "claude-sonnet-5".to_string()]);
assert_eq!(
config.allowed_models.as_deref(),
Some(&["gpt-5".to_string(), "claude-sonnet-5".to_string()][..])
);
assert!(
format!("{config:?}")
.contains("allowed_models: Some([\"gpt-5\", \"claude-sonnet-5\"])")
);

let (wire, _) = config
.into_wire()
.expect("allowed models do not add SDK validation");
let json = serde_json::to_value(&wire).unwrap();
assert_eq!(json["allowedModels"], json!(["gpt-5", "claude-sonnet-5"]));

let (default_wire, _) = ResumeSessionConfig::new(SessionId::from("unrestricted-resume"))
.into_wire()
.expect("default resume config has no duplicate handlers");
let default_json = serde_json::to_value(&default_wire).unwrap();
assert!(default_json.get("allowedModels").is_none());
}

#[test]
fn allowed_models_and_auth_client_id_metadata_url_serialize_together() {
let url = "https://example.com/oauth/client-metadata.json";
let (create_wire, _) = SessionConfig::default()
.with_allowed_models(["gpt-5"])
.with_auth_client_id_metadata_url(url)
.into_wire(None)
.expect("create config has no duplicate handlers");
let (resume_wire, _) =
ResumeSessionConfig::new(SessionId::from("allowed-models-oauth-resume"))
.with_allowed_models(["gpt-5"])
.with_auth_client_id_metadata_url(url)
.into_wire()
.expect("resume config has no duplicate handlers");

for json in [
serde_json::to_value(&create_wire).unwrap(),
serde_json::to_value(&resume_wire).unwrap(),
] {
assert_eq!(json["allowedModels"], json!(["gpt-5"]));
assert_eq!(json["authClientIdMetadataUrl"], url);
}
}

#[test]
fn custom_agents_local_only_serializes_on_create_and_resume() {
let (create_wire, _) = SessionConfig::default()
Expand Down
4 changes: 4 additions & 0 deletions rust/src/wire.rs
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,8 @@ pub(crate) struct SessionCreateWire {
pub session_id: Option<SessionId>,
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(rename = "allowedModels", skip_serializing_if = "Option::is_none")]
pub allowed_models: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub client_name: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
Expand Down Expand Up @@ -212,6 +214,8 @@ pub(crate) struct SessionResumeWire {
pub session_id: SessionId,
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(rename = "allowedModels", skip_serializing_if = "Option::is_none")]
pub allowed_models: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub client_name: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
Expand Down
96 changes: 95 additions & 1 deletion rust/tests/session_test.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,8 @@ use github_copilot_sdk::handler::{
};
use github_copilot_sdk::rpc::{
CanvasProviderInvokeActionRequest, CanvasProviderOpenRequest, CanvasProviderOpenResult,
OpenCanvasInstance, SendAgentMode, SendMode, SendRequest,
ModelSetAllowedModelsRequest, ModelSetAllowedModelsResult, OpenCanvasInstance, SendAgentMode,
SendMode, SendRequest,
};
use github_copilot_sdk::session_events::{
ManagedSettingsResolvedSource, McpOauthRequiredData, ReasoningSummary, SessionLimitsConfig,
Expand Down Expand Up @@ -3223,6 +3224,52 @@ async fn set_model_sends_switch_to_request() {
timeout(TIMEOUT, handle).await.unwrap().unwrap();
}

#[test]
fn model_set_allowed_models_types_serialize_replacement_and_clear() {
let replacement = ModelSetAllowedModelsRequest {
allowed_models: Some(vec!["gpt-5".to_string(), "claude-sonnet-5".to_string()]),
};
assert_eq!(
serde_json::to_value(replacement).unwrap(),
serde_json::json!({
"allowedModels": ["gpt-5", "claude-sonnet-5"]
})
);

assert_eq!(
serde_json::to_value(ModelSetAllowedModelsRequest {
allowed_models: None,
})
.unwrap(),
serde_json::json!({})
);

let result: ModelSetAllowedModelsResult = serde_json::from_value(serde_json::json!({
"allowedModels": ["gpt-5"],
"effectiveAllowedModels": ["gpt-5", "claude-*"],
"fallbackModel": "gpt-5",
"modelId": "gpt-5"
}))
.unwrap();
assert_eq!(
result.allowed_models.as_deref(),
Some(&["gpt-5".to_string()][..])
);
assert_eq!(
result.effective_allowed_models.as_deref(),
Some(&["gpt-5".to_string(), "claude-*".to_string()][..])
);
assert_eq!(result.fallback_model.as_deref(), Some("gpt-5"));
assert_eq!(result.model_id.as_deref(), Some("gpt-5"));

let cleared: ModelSetAllowedModelsResult =
serde_json::from_value(serde_json::json!({})).unwrap();
assert_eq!(cleared.allowed_models, None);
assert_eq!(cleared.effective_allowed_models, None);
assert_eq!(cleared.fallback_model, None);
assert_eq!(cleared.model_id, None);
}

#[tokio::test]
async fn elicitation_returns_typed_result() {
let (session, mut server) =
Expand Down Expand Up @@ -5736,6 +5783,53 @@ async fn rpc_namespace_session_tasks_list_dispatches_correctly() {
assert!(result.tasks.is_empty());
}

#[tokio::test]
async fn rpc_namespace_session_model_set_allowed_models_dispatches_correctly() {
let (session, mut server) = create_session_pair().await;
let session = Arc::new(session);

let s = session.clone();
let handle = tokio::spawn(async move {
s.rpc()
.model()
.set_allowed_models(ModelSetAllowedModelsRequest {
allowed_models: Some(vec!["gpt-5".to_string(), "claude-sonnet-5".to_string()]),
})
.await
});

let request = server.read_request().await;
assert_eq!(request["method"], "session.model.setAllowedModels");
assert_eq!(request["params"]["sessionId"], server.session_id);
assert_eq!(
request["params"]["allowedModels"],
serde_json::json!(["gpt-5", "claude-sonnet-5"])
);
server
.respond(
&request,
serde_json::json!({
"allowedModels": ["gpt-5", "claude-sonnet-5"],
"effectiveAllowedModels": ["gpt-5"],
"fallbackModel": "gpt-5",
"modelId": "gpt-5"
}),
)
.await;

let result = timeout(TIMEOUT, handle).await.unwrap().unwrap().unwrap();
assert_eq!(
result.allowed_models.as_deref(),
Some(&["gpt-5".to_string(), "claude-sonnet-5".to_string()][..])
);
assert_eq!(
result.effective_allowed_models.as_deref(),
Some(&["gpt-5".to_string()][..])
);
assert_eq!(result.fallback_model.as_deref(), Some("gpt-5"));
assert_eq!(result.model_id.as_deref(), Some("gpt-5"));
}

#[tokio::test]
async fn rpc_namespace_client_models_list_dispatches_correctly() {
let (session, mut server) = create_session_pair().await;
Expand Down
Loading
Loading