Skip to content
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,8 @@

import org.junit.jupiter.api.Test;

import com.fasterxml.jackson.databind.node.ObjectNode;

/**
* Coverage tests for generated RPC API classes that are not exercised in
* {@link RpcWrappersTest}. Uses the same {@link StubCaller} pattern to verify
Expand All @@ -24,7 +26,7 @@ class GeneratedRpcApiCoverageTest {
/** A simple stub {@link RpcCaller} that records every call made to it. */
private static final class StubCaller implements RpcCaller {

record Call(String method, Object params) {
record Call(String method, Object params, Class<?> resultType) {
}

final List<Call> calls = new ArrayList<>();
Expand All @@ -33,7 +35,7 @@ record Call(String method, Object params) {
@Override
@SuppressWarnings("unchecked")
public <T> CompletableFuture<T> invoke(String method, Object params, Class<T> resultType) {
calls.add(new Call(method, params));
calls.add(new Call(method, params, resultType));
return CompletableFuture.completedFuture((T) nextResult);
}
}
Expand Down Expand Up @@ -105,6 +107,60 @@ void serverRpc_mcp_config_remove_invokes_correct_method() {
assertSame(params, stub.calls.get(0).params());
}

// ── SessionRpc.model ───────────────────────────────────────────────────

@Test
void sessionRpc_model_setAllowedModels_merges_sessionId_and_returns_result() {
var stub = new StubCaller();
var session = new SessionRpc(stub, "sess-model-policy");
var allowedModels = List.of("gpt-5", "azure/gpt-5");
var expectedResult = new SessionModelSetAllowedModelsResult(allowedModels, List.of("gpt-5"), "gpt-5", "gpt-5");
stub.nextResult = expectedResult;
var request = new SessionModelSetAllowedModelsParams("ignored-session", allowedModels);

assertSame(expectedResult, session.model.setAllowedModels(request).join());

assertEquals(1, stub.calls.size());
var call = stub.calls.get(0);
assertEquals("session.model.setAllowedModels", call.method());
assertEquals(SessionModelSetAllowedModelsResult.class, call.resultType());
var params = assertInstanceOf(ObjectNode.class, call.params());
assertEquals("sess-model-policy", params.get("sessionId").asText());
assertEquals(2, params.get("allowedModels").size());
assertEquals("gpt-5", params.get("allowedModels").get(0).asText());
assertEquals("azure/gpt-5", params.get("allowedModels").get(1).asText());
assertEquals("ignored-session", request.sessionId());
}

@Test
void sessionRpc_model_setAllowedModels_omits_cleared_host_restriction() {
var stub = new StubCaller();
var session = new SessionRpc(stub, "sess-model-policy");

session.model.setAllowedModels(new SessionModelSetAllowedModelsParams(null, null)).join();

assertEquals(1, stub.calls.size());
assertEquals("session.model.setAllowedModels", stub.calls.get(0).method());
var params = assertInstanceOf(ObjectNode.class, stub.calls.get(0).params());
assertEquals("sess-model-policy", params.get("sessionId").asText());
assertFalse(params.has("allowedModels"));
}

@Test
void sessionRpc_model_setAllowedModels_preserves_explicit_empty_list() {
var stub = new StubCaller();
var session = new SessionRpc(stub, "sess-model-policy");

session.model.setAllowedModels(new SessionModelSetAllowedModelsParams(null, List.of())).join();

assertEquals(1, stub.calls.size());
assertEquals("session.model.setAllowedModels", stub.calls.get(0).method());
var params = assertInstanceOf(ObjectNode.class, stub.calls.get(0).params());
assertEquals("sess-model-policy", params.get("sessionId").asText());
assertTrue(params.get("allowedModels").isArray());
assertTrue(params.get("allowedModels").isEmpty());
}

// ── SessionRpc.mode ────────────────────────────────────────────────────

@Test
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -325,6 +325,80 @@ void sessionModelGetCurrentParams_record() {
assertEquals("sess-31", params.sessionId());
}

@Test
void sessionModelSetAllowedModelsParams_round_trips_exact_ids() throws Exception {
var mapper = new ObjectMapper();
var allowedModels = List.of("gpt-5", "azure/gpt-5");
var params = new SessionModelSetAllowedModelsParams("sess-model-policy", allowedModels);

assertEquals("sess-model-policy", params.sessionId());
assertEquals(allowedModels, params.allowedModels());
var json = mapper.readTree(mapper.writeValueAsString(params));
assertEquals("sess-model-policy", json.get("sessionId").asText());
assertEquals(mapper.valueToTree(allowedModels), json.get("allowedModels"));
assertEquals(params, mapper.treeToValue(json, SessionModelSetAllowedModelsParams.class));
}

@Test
void sessionModelSetAllowedModelsParams_distinguishes_clearing_from_empty_list() throws Exception {
var mapper = new ObjectMapper();
var cleared = new SessionModelSetAllowedModelsParams("sess-model-policy", null);

for (var json : List.of("""
{"sessionId":"sess-model-policy"}
""", """
{"sessionId":"sess-model-policy","allowedModels":null}
""")) {
assertEquals(cleared, mapper.readValue(json, SessionModelSetAllowedModelsParams.class));
}
var clearedJson = mapper.readTree(mapper.writeValueAsString(cleared));
assertFalse(clearedJson.has("allowedModels"));

var empty = new SessionModelSetAllowedModelsParams("sess-model-policy", List.of());
var emptyJson = mapper.readTree(mapper.writeValueAsString(empty));
assertTrue(emptyJson.get("allowedModels").isArray());
assertTrue(emptyJson.get("allowedModels").isEmpty());
assertEquals(empty, mapper.treeToValue(emptyJson, SessionModelSetAllowedModelsParams.class));
}

@Test
void sessionModelSetAllowedModelsResult_round_trips_policy_and_selection() throws Exception {
var mapper = new ObjectMapper();
var json = """
{"allowedModels":["gpt-5","azure/gpt-5"],"effectiveAllowedModels":["gpt-5"],
"fallbackModel":"gpt-5","modelId":"gpt-5"}
""";

var result = mapper.readValue(json, SessionModelSetAllowedModelsResult.class);

assertEquals(List.of("gpt-5", "azure/gpt-5"), result.allowedModels());
assertEquals(List.of("gpt-5"), result.effectiveAllowedModels());
assertEquals("gpt-5", result.fallbackModel());
assertEquals("gpt-5", result.modelId());
assertEquals(mapper.readTree(json), mapper.valueToTree(result));
}

@Test
void sessionModelSetAllowedModelsResult_all_fields_are_optional() throws Exception {
var mapper = new ObjectMapper();
var empty = mapper.readValue("{}", SessionModelSetAllowedModelsResult.class);
assertNull(empty.allowedModels());
assertNull(empty.effectiveAllowedModels());
assertNull(empty.fallbackModel());
assertNull(empty.modelId());

for (var json : List.of("{}", """
{"modelId":"gpt-5"}
""", """
{"allowedModels":["gpt-5"]}
""", """
{"effectiveAllowedModels":[],"fallbackModel":"gpt-5"}
""")) {
var result = mapper.readValue(json, SessionModelSetAllowedModelsResult.class);
assertEquals(mapper.readTree(json), mapper.valueToTree(result));
}
}

@Test
void sessionModelSwitchToParams_record() {
var params = new SessionModelSwitchToParams("sess-32", "claude-sonnet-5", null, "high", null, null, null, null,
Expand Down
82 changes: 82 additions & 0 deletions rust/src/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1944,6 +1944,11 @@ pub struct SessionConfig {
pub session_id: Option<SessionId>,
/// Model to use (e.g. `"gpt-4"`, `"claude-sonnet-4"`).
pub model: Option<String>,
/// Exact model IDs this session may use. When unset, the host imposes no
/// model restriction. The runtime validates configured IDs, rejects an
/// explicit empty list, and intersects the list with applicable model
/// policies.
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 +2307,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 +2455,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 +2627,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 +2852,19 @@ impl SessionConfig {
self
}

/// Restrict this session to the provided exact model IDs.
///
/// Passing an empty iterator sends an explicit empty list, which the
/// runtime rejects.
pub fn with_allowed_models<I, S>(mut self, models: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
self.allowed_models = Some(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 +3423,11 @@ 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 IDs the resumed session may use. When unset, the host
/// imposes no model restriction. The runtime validates configured IDs,
/// rejects an explicit empty list, and intersects the list with applicable
/// model policies.
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 +3698,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 +3889,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 +3999,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 +4195,19 @@ impl ResumeSessionConfig {
self
}

/// Restrict the resumed session to the provided exact model IDs.
///
/// Passing an empty iterator sends an explicit empty list, which the
/// runtime rejects.
pub fn with_allowed_models<I, S>(mut self, models: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
self.allowed_models = Some(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 @@ -6493,6 +6535,7 @@ mod tests {
fn session_config_default_wire_flags_off_without_handlers() {
let cfg = SessionConfig::default();
assert_eq!(cfg.mcp_oauth_token_storage, None);
assert_eq!(cfg.allowed_models, None);
// Wire flags are derived from handler presence at create_session
// time, not stored on the config. With no handlers installed, every
// request_* flag should serialize as false.
Expand All @@ -6508,12 +6551,14 @@ mod tests {
assert!(!wire.request_mcp_apps);
let json = serde_json::to_value(&wire).unwrap();
assert!(json.get("askUserVariant").is_none());
assert!(json.get("allowedModels").is_none());
}

#[test]
fn resume_session_config_new_wire_flags_off_without_handlers() {
let cfg = ResumeSessionConfig::new(SessionId::from("resume-flags"));
assert_eq!(cfg.mcp_oauth_token_storage, None);
assert_eq!(cfg.allowed_models, None);
let (wire, _runtime) = cfg
.into_wire()
.expect("default resume config has no duplicate handlers");
Expand All @@ -6526,6 +6571,43 @@ mod tests {
assert!(!wire.request_mcp_apps);
let json = serde_json::to_value(&wire).unwrap();
assert!(json.get("askUserVariant").is_none());
assert!(json.get("allowedModels").is_none());
}

#[test]
fn session_configs_build_debug_and_serialize_allowed_models() {
let create = SessionConfig::default().with_allowed_models(["gpt-5.4", "claude-sonnet-4"]);
assert_eq!(
create.allowed_models.as_deref(),
Some(&["gpt-5.4".to_string(), "claude-sonnet-4".to_string()][..])
);
assert!(format!("{create:?}").contains("allowed_models"));

let (create_wire, _) = create
.into_wire(Some(SessionId::from("create-allowed-models")))
.expect("allowed model config has no duplicate handlers");
let create_json = serde_json::to_value(&create_wire).unwrap();
assert_eq!(
create_json["allowedModels"],
json!(["gpt-5.4", "claude-sonnet-4"])
);

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

let (resume_wire, _) = resume
.into_wire()
.expect("resume allowed model config has no duplicate handlers");
let resume_json = serde_json::to_value(&resume_wire).unwrap();
assert_eq!(
resume_json["allowedModels"],
json!(["gpt-5.4", "gpt-5-mini"])
);
}

#[test]
Expand Down
4 changes: 4 additions & 0 deletions rust/src/wire.rs
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,8 @@ pub(crate) struct SessionCreateWire {
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(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")]
pub reasoning_effort: Option<String>,
Expand Down Expand Up @@ -213,6 +215,8 @@ pub(crate) struct SessionResumeWire {
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(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")]
pub reasoning_effort: Option<String>,
Expand Down
Loading
Loading