Skip to content
Closed
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
12 changes: 6 additions & 6 deletions python/copilot/generated/rpc.py

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

51 changes: 51 additions & 0 deletions python/test_rpc_generated.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,15 @@

import pytest

from copilot.generated.rpc import _load_SessionListEntry
from copilot.rpc import (
CommandsApi,
CommandsInvokeRequest,
LocalSessionMetadataValue,
QueuedCommandHandled,
QueuedCommandNotHandled,
RemoteSessionMetadataValue,
SessionList,
SlashCommandTextResult,
)

Expand All @@ -22,3 +28,48 @@ async def test_commands_invoke_deserializes_slash_command_result():
assert isinstance(result, SlashCommandTextResult)
assert result.text == "hello"
assert result.markdown is True


def test_session_list_entry_decodes_boolean_is_remote_discriminator():
local_payload = {
"sessionId": "example-local",
"startTime": "2026-07-26T10:00:00.000Z",
"modifiedTime": "2026-07-26T10:05:00.000Z",
"isRemote": False,
}
remote_payload = {
"sessionId": "example-remote",
"startTime": "2026-07-26T10:00:00.000Z",
"modifiedTime": "2026-07-26T10:05:00.000Z",
"isRemote": True,
"remoteSessionIds": ["rs-1"],
"repository": {
"owner": "github",
"name": "copilot-sdk",
"branch": "main",
},
}

local_entry = _load_SessionListEntry(local_payload)
remote_entry = _load_SessionListEntry(remote_payload)

assert isinstance(local_entry, LocalSessionMetadataValue)
assert local_entry.session_id == "example-local"
assert isinstance(remote_entry, RemoteSessionMetadataValue)
assert remote_entry.session_id == "example-remote"

session_list = SessionList.from_dict({"sessions": [local_payload, remote_payload]})
assert len(session_list.sessions) == 2
assert isinstance(session_list.sessions[0], LocalSessionMetadataValue)
assert isinstance(session_list.sessions[1], RemoteSessionMetadataValue)


def test_queued_command_result_round_trips_boolean_handled_discriminator():
handled = QueuedCommandHandled(stop_processing_queue=True)
not_handled = QueuedCommandNotHandled()

assert handled.to_dict() == {"handled": True, "stopProcessingQueue": True}
assert not_handled.to_dict() == {"handled": False}

assert QueuedCommandHandled.from_dict({"handled": True, "stopProcessingQueue": True}) == handled
assert QueuedCommandNotHandled.from_dict({"handled": False}) == not_handled
60 changes: 49 additions & 11 deletions scripts/codegen/python.ts
Original file line number Diff line number Diff line change
Expand Up @@ -272,6 +272,33 @@ function postProcessExternalUnionAliasesForPython(code: string, aliases: Map<str
return code.replace(/\n{3,}/g, "\n\n");
}

type PyDiscriminatorValue = string | boolean;

function pyDiscriminatorConstValue(schema: JSONSchema7): PyDiscriminatorValue | undefined {
if (typeof schema.const === "string" || typeof schema.const === "boolean") {
return schema.const;
}
return undefined;
}

function pyDiscriminatorMatchPattern(value: PyDiscriminatorValue): string {
if (typeof value === "boolean") {
return value ? "True" : "False";
}
return JSON.stringify(value);
}

function pyDiscriminatorLiteral(value: PyDiscriminatorValue): string {
if (typeof value === "boolean") {
return value ? "True" : "False";
}
return JSON.stringify(value);
}

function pyDiscriminatorClassVarType(value: PyDiscriminatorValue): "str" | "bool" {
return typeof value === "boolean" ? "bool" : "str";
}

/**
* Replace flat-merged dataclasses emitted by quicktype for $ref-based
* discriminated unions with proper Python unions: a `Name = VariantA | ...`
Expand All @@ -293,7 +320,7 @@ function postProcessExternalUnionAliasesForPython(code: string, aliases: Map<str
interface ResolvedRefBasedUnion {
aliasName: string;
discriminatorProp: string;
dispatch: Array<{ value: string; typeName: string }>;
dispatch: Array<{ value: PyDiscriminatorValue; typeName: string }>;
}
function postProcessRefBasedDiscriminatedUnionsForPython(
code: string,
Expand All @@ -304,7 +331,7 @@ function postProcessRefBasedDiscriminatedUnionsForPython(
aliasName: string;
variantNames: string[];
discriminatorProp: string;
dispatch: Array<{ value: string; typeName: string }>;
dispatch: Array<{ value: PyDiscriminatorValue; typeName: string }>;
description: string | undefined;
}
const unions: UnionInfo[] = [];
Expand Down Expand Up @@ -333,8 +360,12 @@ function postProcessRefBasedDiscriminatedUnionsForPython(
const discProp = (resolvedVariants[i].properties as Record<string, JSONSchema7>)[
discriminator.property
];
const discValue = pyDiscriminatorConstValue(discProp);
if (discValue === undefined) {
throw new Error(`Missing discriminator const on ${variantRefNames[i]}`);
}
return {
value: String(discProp.const),
value: discValue,
typeName: toPascalCase(variantRefNames[i]),
};
});
Expand Down Expand Up @@ -387,7 +418,7 @@ function postProcessRefBasedDiscriminatedUnionsForPython(
for (const union of unions) {
const actualAliasName = resolveActualName(union.aliasName);
const actualVariantNames: string[] = [];
const actualDispatch: Array<{ value: string; typeName: string }> = [];
const actualDispatch: Array<{ value: PyDiscriminatorValue; typeName: string }> = [];
let allResolved = true;
for (let i = 0; i < union.variantNames.length; i++) {
const actual = resolveActualName(union.variantNames[i]);
Expand Down Expand Up @@ -450,7 +481,9 @@ function postProcessRefBasedDiscriminatedUnionsForPython(
dispatcherLines.push(` kind = obj.get(${JSON.stringify(union.discriminatorProp)})`);
dispatcherLines.push(` match kind:`);
for (const m of actualDispatch) {
dispatcherLines.push(` case ${JSON.stringify(m.value)}: return ${m.typeName}.from_dict(obj)`);
dispatcherLines.push(
` case ${pyDiscriminatorMatchPattern(m.value)}: return ${m.typeName}.from_dict(obj)`
);
}
dispatcherLines.push(
` case _: raise ValueError(f"Unknown ${actualAliasName} ${union.discriminatorProp}: {kind!r}")`
Expand Down Expand Up @@ -500,7 +533,7 @@ function postProcessDiscriminatorDefaultsForPython(
unions: ResolvedRefBasedUnion[]
): string {
// Build variant lookup: variant class name → { prop, value }.
const variantInfo = new Map<string, { prop: string; value: string }>();
const variantInfo = new Map<string, { prop: string; value: PyDiscriminatorValue }>();
for (const union of unions) {
for (const d of union.dispatch) {
// First-wins; multiple unions referencing the same variant share a
Expand Down Expand Up @@ -571,9 +604,10 @@ function postProcessDiscriminatorDefaultsForPython(
continue;
}
const fieldIndent = (block[fieldIdx].match(/^(\s+)/) ?? ["", ""])[1];
const literal = JSON.stringify(info.value);
const literal = pyDiscriminatorLiteral(info.value);
const classVarType = pyDiscriminatorClassVarType(info.value);
// Replace the field with a class-level constant.
block[fieldIdx] = `${fieldIndent}${info.prop}: ClassVar[str] = ${literal}`;
block[fieldIdx] = `${fieldIndent}${info.prop}: ClassVar[${classVarType}] = ${literal}`;
usedClassVar = true;

// Drop any field-trailing docstring lines that immediately followed the
Expand Down Expand Up @@ -1590,7 +1624,7 @@ function tryEmitPyRefBasedDiscriminatedUnion(
if (!discriminator) return undefined;

const variantTypeNames: string[] = [];
const dispatch: Array<{ value: string; typeName: string }> = [];
const dispatch: Array<{ value: PyDiscriminatorValue; typeName: string }> = [];
for (let i = 0; i < variants.length; i++) {
const variantTypeName = toPascalCase(variantRefNames[i]);
const variantSchema = resolveObjectSchema(variants[i], ctx.definitions);
Expand All @@ -1599,7 +1633,11 @@ function tryEmitPyRefBasedDiscriminatedUnion(
}
variantTypeNames.push(variantTypeName);
const discProp = resolvedVariants[i].properties?.[discriminator.property] as JSONSchema7;
dispatch.push({ value: String(discProp.const), typeName: variantTypeName });
const discValue = pyDiscriminatorConstValue(discProp);
if (discValue === undefined) {
throw new Error(`Missing discriminator const on ${variantRefNames[i]}`);
}
dispatch.push({ value: discValue, typeName: variantTypeName });
}

if (!ctx.aliasesByName.has(aliasName)) {
Expand Down Expand Up @@ -1627,7 +1665,7 @@ function tryEmitPyRefBasedDiscriminatedUnion(
lines.push(` match kind:`);
for (const m of dispatch) {
lines.push(
` case ${JSON.stringify(m.value)}: return ${m.typeName}.from_dict(obj)`
` case ${pyDiscriminatorMatchPattern(m.value)}: return ${m.typeName}.from_dict(obj)`
);
}
lines.push(
Expand Down