diff --git a/rust/src/generated/api_types.rs b/rust/src/generated/api_types.rs index c93f34cf64..18d4c7ed58 100644 --- a/rust/src/generated/api_types.rs +++ b/rust/src/generated/api_types.rs @@ -8484,7 +8484,7 @@ pub struct ModelSetReasoningEffortResult { pub reasoning_effort: String, } -/// Optional GitHub token used to list models for a specific user instead of the global auth context. +/// Optional GitHub token and working directory used to resolve available models. /// ///
/// @@ -8495,6 +8495,9 @@ pub struct ModelSetReasoningEffortResult { #[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct ModelsListRequest { + /// Working directory used to apply repository model policy. When omitted, model availability is account-global. + #[serde(skip_serializing_if = "Option::is_none")] + pub cwd: Option, /// GitHub token for per-user model listing. When provided, resolves this token to determine the user's Copilot plan and available models instead of using the global auth. #[serde(skip_serializing_if = "Option::is_none")] pub git_hub_token: Option, diff --git a/rust/tests/api_types_test.rs b/rust/tests/api_types_test.rs index 9b86b1367a..c57272c2fc 100644 --- a/rust/tests/api_types_test.rs +++ b/rust/tests/api_types_test.rs @@ -5,7 +5,8 @@ use github_copilot_sdk::rpc::{ Extension, ExtensionList, ExtensionSource, ExtensionStatus, ExtensionsDisableRequest, - ExtensionsEnableRequest, FleetStartRequest, FleetStartResult, TasksStartAgentRequest, + ExtensionsEnableRequest, FleetStartRequest, FleetStartResult, ModelsListRequest, + TasksStartAgentRequest, }; use github_copilot_sdk::session_events::{PermissionRequest, PermissionRequestedData}; @@ -104,6 +105,23 @@ fn permission_event_exposes_managed_approval_required() { assert_eq!(request.managed_approval_required, Some(true)); } +#[test] +fn models_list_request_serializes_repository_cwd() { + let request = ModelsListRequest { + cwd: Some("/workspace/repository".to_string()), + git_hub_token: None, + }; + + assert_eq!( + serde_json::to_value(request).unwrap(), + serde_json::json!({ "cwd": "/workspace/repository" }) + ); + assert_eq!( + serde_json::to_value(ModelsListRequest::default()).unwrap(), + serde_json::json!({}) + ); +} + fn running_extension(id: &str, name: &str) -> Extension { Extension { id: id.to_string(), diff --git a/rust/tests/session_test.rs b/rust/tests/session_test.rs index 727911081d..bb74f4cae8 100644 --- a/rust/tests/session_test.rs +++ b/rust/tests/session_test.rs @@ -14,7 +14,7 @@ use github_copilot_sdk::handler::{ }; use github_copilot_sdk::rpc::{ CanvasProviderInvokeActionRequest, CanvasProviderOpenRequest, CanvasProviderOpenResult, - OpenCanvasInstance, + ModelsListRequest, OpenCanvasInstance, }; use github_copilot_sdk::session_events::{ ManagedSettingsResolvedSource, McpOauthRequiredData, ReasoningSummary, SessionLimitsConfig, @@ -4044,6 +4044,39 @@ async fn rpc_namespace_client_models_list_dispatches_correctly() { assert!(result.models.is_empty()); } +#[tokio::test] +async fn rpc_namespace_client_models_list_sends_repository_cwd() { + let (session, mut server) = create_session_pair().await; + let session = Arc::new(session); + + let client = session.client().clone(); + let handle = tokio::spawn(async move { + client + .rpc() + .models() + .list_with_params(ModelsListRequest { + cwd: Some("/workspace/repository".to_string()), + git_hub_token: None, + }) + .await + }); + + let request = server.read_request().await; + assert_eq!(request["method"], "models.list"); + assert_eq!( + request["params"], + serde_json::json!({ + "cwd": "/workspace/repository" + }) + ); + server + .respond(&request, serde_json::json!({ "models": [] })) + .await; + + let result = timeout(TIMEOUT, handle).await.unwrap().unwrap().unwrap(); + assert!(result.models.is_empty()); +} + #[tokio::test] async fn client_stop_sends_session_destroy_for_each_active_session() { // One client, two registered sessions. Client::stop must send diff --git a/scripts/codegen/rust.ts b/scripts/codegen/rust.ts index 4090318f05..1375af4eb1 100644 --- a/scripts/codegen/rust.ts +++ b/scripts/codegen/rust.ts @@ -17,6 +17,7 @@ import { fileURLToPath } from "url"; import { promisify } from "util"; import type { JSONSchema7, JSONSchema7Definition } from "json-schema"; import { + addCwdToModelsListRequest, addManagedApprovalRequiredToPermissionRequests, type ApiSchema, type DefinitionCollections, @@ -2219,7 +2220,7 @@ async function generate(): Promise { ); const apiSchema = propagateInternalVisibility( postProcessSchema( - stripBooleanLiterals(apiRaw) as JSONSchema7, + stripBooleanLiterals(addCwdToModelsListRequest(apiRaw)) as JSONSchema7, ), ) as unknown as ApiSchema; diff --git a/scripts/codegen/utils.ts b/scripts/codegen/utils.ts index 42e78b9a07..089065ed34 100644 --- a/scripts/codegen/utils.ts +++ b/scripts/codegen/utils.ts @@ -496,6 +496,46 @@ export function addManagedApprovalRequiredToPermissionRequests(schema: T): T { + const cloned = cloneSchemaForCodegen(schema); + const property: JSONSchema7 = { + description: + "Working directory used to apply repository model policy. When omitted, model availability is account-global.", + type: ["string", "null"], + }; + (property as Record)["x-copilot-sdk-append-last"] = true; + + for (const definitions of [cloned.definitions, cloned.$defs]) { + if (!definitions) continue; + const definition = definitions.ModelsListRequest; + if (!definition || typeof definition !== "object") continue; + const requestDefinition = definition as JSONSchema7; + const objectDefinition = [ + requestDefinition, + ...(requestDefinition.anyOf ?? []), + ...(requestDefinition.oneOf ?? []), + ].find( + (candidate): candidate is JSONSchema7 => + typeof candidate === "object" && + candidate !== null && + (candidate.type === "object" || candidate.properties !== undefined), + ); + if (!objectDefinition) continue; + if (objectDefinition.properties?.cwd) continue; + requestDefinition.description = + "Optional GitHub token and working directory used to resolve available models."; + objectDefinition.properties = { + ...objectDefinition.properties, + cwd: cloneSchemaForCodegen(property), + }; + } + + return cloned; +} + export function getEnumValueDescriptions(schema: JSONSchema7 | null | undefined): EnumValueDescriptions | undefined { if (!schema || typeof schema !== "object") return undefined;