From 209c756f132e92f4cc11abd5d52b738a19d893e9 Mon Sep 17 00:00:00 2001 From: Jake Bailey <5341706+jakebailey@users.noreply.github.com> Date: Thu, 20 Aug 2026 14:27:57 -0700 Subject: [PATCH 1/3] Optimize LSP discriminated union decoding --- .../lsp/lsproto/_generate/generate.mts | 97 ++++++-- tsc/internal/lsp/lsproto/lsp.go | 30 --- tsc/internal/lsp/lsproto/lsp_generated.go | 207 +++++++++++------- .../lsp/lsproto/lsp_json_benchmark_test.go | 87 ++++++++ tsc/internal/lsp/lsproto/lsp_json_test.go | 27 +++ tsc/internal/lsp/lsproto/structcodec.go | 168 ++++++++++++-- 6 files changed, 465 insertions(+), 151 deletions(-) create mode 100644 tsc/internal/lsp/lsproto/lsp_json_benchmark_test.go diff --git a/tsc/internal/lsp/lsproto/_generate/generate.mts b/tsc/internal/lsp/lsproto/_generate/generate.mts index 0fb164a65cf62..fe5bb9230c0b5 100755 --- a/tsc/internal/lsp/lsproto/_generate/generate.mts +++ b/tsc/internal/lsp/lsproto/_generate/generate.mts @@ -2196,6 +2196,47 @@ function generateCode() { return exhaustive; } + /** + * Generate streaming discriminator dispatch for unions with at most one + * unmapped fallback arm. Fields after the discriminator decode directly + * from the active decoder; fields before it are replayed by the helper. + */ + function generateStreamingDiscriminatorDispatch( + name: string, + disc: NonNullable>, + indent: string, + ) { + writeLine(`${indent}target, err := unmarshalDiscriminatedStruct(dec, "${name}", ${JSON.stringify(disc.fieldName)}, func(value json.Value) any {`); + writeLine(`${indent}\tswitch string(value) {`); + for (const [value, entry] of disc.mapping) { + writeLine(`${indent}\tcase \`"${value}"\`:`); + writeLine(`${indent}\t\treturn new(${entry.typeName})`); + } + writeLine(`${indent}\tdefault:`); + if (disc.unmapped.length === 1) { + writeLine(`${indent}\t\treturn new(${disc.unmapped[0].typeName})`); + } + else { + writeLine(`${indent}\t\treturn nil`); + } + writeLine(`${indent}\t}`); + writeLine(`${indent}})`); + writeLine(`${indent}if err != nil {`); + writeLine(`${indent}\treturn err`); + writeLine(`${indent}}`); + writeLine(`${indent}switch target := target.(type) {`); + for (const entry of disc.mapping.values()) { + writeLine(`${indent}case *${entry.typeName}:`); + writeLine(`${indent}\to.${entry.fieldName} = target`); + } + for (const entry of disc.unmapped) { + writeLine(`${indent}case *${entry.typeName}:`); + writeLine(`${indent}\to.${entry.fieldName} = target`); + } + writeLine(`${indent}}`); + writeLine(`${indent}return nil`); + } + /** * Generate try-each fallback code for unmapped entries, chaining into * presence dispatch if possible before falling back to raw try-each. @@ -3381,17 +3422,24 @@ function generateCode() { } } else { - // Ambiguous: buffer and dispatch - writeLine(`\t\tdata, err := dec.ReadValue()`); - writeLine(`\t\tif err != nil {`); - writeLine(`\t\t\treturn err`); - writeLine(`\t\t}`); let exhaustive = false; const disc = findDiscriminatorField(entries); - if (disc) { - exhaustive = generateDiscriminatorDispatch(disc, "\t\t"); + if (disc && disc.unmapped.length <= 1) { + generateStreamingDiscriminatorDispatch(name, disc, "\t\t"); + exhaustive = true; } else { + // Ambiguous non-discriminated objects need the complete + // value for presence checks or speculative decoding. + writeLine(`\t\tdata, err := dec.ReadValue()`); + writeLine(`\t\tif err != nil {`); + writeLine(`\t\t\treturn err`); + writeLine(`\t\t}`); + } + if (disc && disc.unmapped.length > 1) { + exhaustive = generateDiscriminatorDispatch(disc, "\t\t"); + } + else if (!disc) { const pres = findPresenceDiscriminator(entries); if (pres) { exhaustive = generatePresenceDispatch(pres, "\t\t"); @@ -3417,25 +3465,30 @@ function generateCode() { writeLine(`\t}`); } else { - // Fallback: unknown kinds present (e.g. `any`), use ReadValue + try-each. - writeLine("\tdata, err := dec.ReadValue()"); - writeLine("\tif err != nil {"); - writeLine("\t\treturn err"); - writeLine("\t}"); - - if (unionContainedNull) { - writeLine(`\tif string(data) == "null" {`); - writeLine(`\t\treturn nil`); - writeLine(`\t}`); - writeLine(""); - } - + // Fallback for unknown kinds (e.g. `any`). Discriminated object + // unions can still stream; other unions use ReadValue + try-each. let exhaustive = false; const disc = findDiscriminatorField(fieldEntries); - if (disc) { - exhaustive = generateDiscriminatorDispatch(disc, "\t"); + if (disc && disc.unmapped.length <= 1) { + generateStreamingDiscriminatorDispatch(name, disc, "\t"); + exhaustive = true; } else { + writeLine("\tdata, err := dec.ReadValue()"); + writeLine("\tif err != nil {"); + writeLine("\t\treturn err"); + writeLine("\t}"); + if (unionContainedNull) { + writeLine(`\tif string(data) == "null" {`); + writeLine(`\t\treturn nil`); + writeLine(`\t}`); + writeLine(""); + } + } + if (disc && disc.unmapped.length > 1) { + exhaustive = generateDiscriminatorDispatch(disc, "\t"); + } + else if (!disc) { const pres = findPresenceDiscriminator(fieldEntries); if (pres) { exhaustive = generatePresenceDispatch(pres, "\t"); diff --git a/tsc/internal/lsp/lsproto/lsp.go b/tsc/internal/lsp/lsproto/lsp.go index ae6c15b589ba1..9197f1c1dea1e 100644 --- a/tsc/internal/lsp/lsproto/lsp.go +++ b/tsc/internal/lsp/lsproto/lsp.go @@ -125,36 +125,6 @@ func jsonKeyCheck(name []byte, key string) bool { return len(name) == len(key)+2 && name[0] == '"' && string(name[1:len(name)-1]) == key } -// jsonObjectRawField scans the top-level keys of a JSON object looking for the -// given field name, and returns its raw JSON value (e.g. `"full"` with quotes). -// Returns nil if the field is not found. -func jsonObjectRawField(data []byte, field string) json.Value { - dec := json.NewDecoder(bytes.NewBuffer(data)) - if dec.PeekKind() != '{' { - return nil - } - if _, err := dec.ReadToken(); err != nil { - return nil - } - for dec.PeekKind() != '}' { - name, err := dec.ReadValue() - if err != nil { - return nil - } - if jsonKeyCheck(name, field) { - val, err := dec.ReadValue() - if err != nil { - return nil - } - return val - } - if err := dec.SkipValue(); err != nil { - return nil - } - } - return nil -} - // jsonObjectHasKey scans the top-level keys of a JSON object looking for any of the // given keys. Returns the index of the first key found, or -1 if none match. // Bails early on first match without decoding any values. diff --git a/tsc/internal/lsp/lsproto/lsp_generated.go b/tsc/internal/lsp/lsproto/lsp_generated.go index 31e2c37a7f617..9c631db1aa18c 100644 --- a/tsc/internal/lsp/lsproto/lsp_generated.go +++ b/tsc/internal/lsp/lsproto/lsp_generated.go @@ -11891,24 +11891,32 @@ var _ json.UnmarshalerFrom = (*TextDocumentEditOrCreateFileOrRenameFileOrDeleteF func (o *TextDocumentEditOrCreateFileOrRenameFileOrDeleteFile) UnmarshalJSONFrom(dec *json.Decoder) error { *o = TextDocumentEditOrCreateFileOrRenameFileOrDeleteFile{} - data, err := dec.ReadValue() + target, err := unmarshalDiscriminatedStruct(dec, "TextDocumentEditOrCreateFileOrRenameFileOrDeleteFile", "kind", func(value json.Value) any { + switch string(value) { + case `"rename"`: + return new(RenameFile) + case `"create"`: + return new(CreateFile) + case `"delete"`: + return new(DeleteFile) + default: + return new(TextDocumentEdit) + } + }) if err != nil { return err } - switch string(jsonObjectRawField(data, "kind")) { - case `"rename"`: - o.RenameFile = new(RenameFile) - return json.Unmarshal(data, o.RenameFile) - case `"create"`: - o.CreateFile = new(CreateFile) - return json.Unmarshal(data, o.CreateFile) - case `"delete"`: - o.DeleteFile = new(DeleteFile) - return json.Unmarshal(data, o.DeleteFile) - default: - o.TextDocumentEdit = new(TextDocumentEdit) - return json.Unmarshal(data, o.TextDocumentEdit) + switch target := target.(type) { + case *RenameFile: + o.RenameFile = target + case *CreateFile: + o.CreateFile = target + case *DeleteFile: + o.DeleteFile = target + case *TextDocumentEdit: + o.TextDocumentEdit = target } + return nil } type StringOrInlayHintLabelParts struct { @@ -11983,19 +11991,26 @@ var _ json.UnmarshalerFrom = (*WorkspaceFullDocumentDiagnosticReportOrUnchangedD func (o *WorkspaceFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport) UnmarshalJSONFrom(dec *json.Decoder) error { *o = WorkspaceFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport{} - data, err := dec.ReadValue() + target, err := unmarshalDiscriminatedStruct(dec, "WorkspaceFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport", "kind", func(value json.Value) any { + switch string(value) { + case `"full"`: + return new(WorkspaceFullDocumentDiagnosticReport) + case `"unchanged"`: + return new(WorkspaceUnchangedDocumentDiagnosticReport) + default: + return nil + } + }) if err != nil { return err } - switch string(jsonObjectRawField(data, "kind")) { - case `"full"`: - o.FullDocumentDiagnosticReport = new(WorkspaceFullDocumentDiagnosticReport) - return json.Unmarshal(data, o.FullDocumentDiagnosticReport) - case `"unchanged"`: - o.UnchangedDocumentDiagnosticReport = new(WorkspaceUnchangedDocumentDiagnosticReport) - return json.Unmarshal(data, o.UnchangedDocumentDiagnosticReport) + switch target := target.(type) { + case *WorkspaceFullDocumentDiagnosticReport: + o.FullDocumentDiagnosticReport = target + case *WorkspaceUnchangedDocumentDiagnosticReport: + o.UnchangedDocumentDiagnosticReport = target } - return errInvalidValue("WorkspaceFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport", data) + return nil } type StringOrStringValue struct { @@ -12368,22 +12383,30 @@ var _ json.UnmarshalerFrom = (*WorkDoneProgressBeginOrReportOrEnd)(nil) func (o *WorkDoneProgressBeginOrReportOrEnd) UnmarshalJSONFrom(dec *json.Decoder) error { *o = WorkDoneProgressBeginOrReportOrEnd{} - data, err := dec.ReadValue() + target, err := unmarshalDiscriminatedStruct(dec, "WorkDoneProgressBeginOrReportOrEnd", "kind", func(value json.Value) any { + switch string(value) { + case `"begin"`: + return new(WorkDoneProgressBegin) + case `"report"`: + return new(WorkDoneProgressReport) + case `"end"`: + return new(WorkDoneProgressEnd) + default: + return nil + } + }) if err != nil { return err } - switch string(jsonObjectRawField(data, "kind")) { - case `"begin"`: - o.Begin = new(WorkDoneProgressBegin) - return json.Unmarshal(data, o.Begin) - case `"report"`: - o.Report = new(WorkDoneProgressReport) - return json.Unmarshal(data, o.Report) - case `"end"`: - o.End = new(WorkDoneProgressEnd) - return json.Unmarshal(data, o.End) + switch target := target.(type) { + case *WorkDoneProgressBegin: + o.Begin = target + case *WorkDoneProgressReport: + o.Report = target + case *WorkDoneProgressEnd: + o.End = target } - return errInvalidValue("WorkDoneProgressBeginOrReportOrEnd", data) + return nil } type TextEditOrAnnotatedTextEditOrSnippetTextEdit struct { @@ -12436,19 +12459,26 @@ var _ json.UnmarshalerFrom = (*FullDocumentDiagnosticReportOrUnchangedDocumentDi func (o *FullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport) UnmarshalJSONFrom(dec *json.Decoder) error { *o = FullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport{} - data, err := dec.ReadValue() + target, err := unmarshalDiscriminatedStruct(dec, "FullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport", "kind", func(value json.Value) any { + switch string(value) { + case `"full"`: + return new(FullDocumentDiagnosticReport) + case `"unchanged"`: + return new(UnchangedDocumentDiagnosticReport) + default: + return nil + } + }) if err != nil { return err } - switch string(jsonObjectRawField(data, "kind")) { - case `"full"`: - o.FullDocumentDiagnosticReport = new(FullDocumentDiagnosticReport) - return json.Unmarshal(data, o.FullDocumentDiagnosticReport) - case `"unchanged"`: - o.UnchangedDocumentDiagnosticReport = new(UnchangedDocumentDiagnosticReport) - return json.Unmarshal(data, o.UnchangedDocumentDiagnosticReport) + switch target := target.(type) { + case *FullDocumentDiagnosticReport: + o.FullDocumentDiagnosticReport = target + case *UnchangedDocumentDiagnosticReport: + o.UnchangedDocumentDiagnosticReport = target } - return errInvalidValue("FullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport", data) + return nil } type TextDocumentSyncOptionsOrKind struct { @@ -13590,22 +13620,30 @@ var _ json.UnmarshalerFrom = (*VSImageElementOrClassifiedTextElementOrContainerE func (o *VSImageElementOrClassifiedTextElementOrContainerElement) UnmarshalJSONFrom(dec *json.Decoder) error { *o = VSImageElementOrClassifiedTextElementOrContainerElement{} - data, err := dec.ReadValue() + target, err := unmarshalDiscriminatedStruct(dec, "VSImageElementOrClassifiedTextElementOrContainerElement", "_vs_type", func(value json.Value) any { + switch string(value) { + case `"ContainerElement"`: + return new(VSContainerElement) + case `"ImageElement"`: + return new(VSImageElement) + case `"ClassifiedTextElement"`: + return new(VSClassifiedTextElement) + default: + return nil + } + }) if err != nil { return err } - switch string(jsonObjectRawField(data, "_vs_type")) { - case `"ContainerElement"`: - o.ContainerElement = new(VSContainerElement) - return json.Unmarshal(data, o.ContainerElement) - case `"ImageElement"`: - o.ImageElement = new(VSImageElement) - return json.Unmarshal(data, o.ImageElement) - case `"ClassifiedTextElement"`: - o.ClassifiedTextElement = new(VSClassifiedTextElement) - return json.Unmarshal(data, o.ClassifiedTextElement) + switch target := target.(type) { + case *VSContainerElement: + o.ContainerElement = target + case *VSImageElement: + o.ImageElement = target + case *VSClassifiedTextElement: + o.ClassifiedTextElement = target } - return errInvalidValue("VSImageElementOrClassifiedTextElementOrContainerElement", data) + return nil } type LocationOrLocationsOrDefinitionLinksOrNull struct { @@ -14085,19 +14123,26 @@ var _ json.UnmarshalerFrom = (*RelatedFullDocumentDiagnosticReportOrUnchangedDoc func (o *RelatedFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport) UnmarshalJSONFrom(dec *json.Decoder) error { *o = RelatedFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport{} - data, err := dec.ReadValue() + target, err := unmarshalDiscriminatedStruct(dec, "RelatedFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport", "kind", func(value json.Value) any { + switch string(value) { + case `"full"`: + return new(RelatedFullDocumentDiagnosticReport) + case `"unchanged"`: + return new(RelatedUnchangedDocumentDiagnosticReport) + default: + return nil + } + }) if err != nil { return err } - switch string(jsonObjectRawField(data, "kind")) { - case `"full"`: - o.FullDocumentDiagnosticReport = new(RelatedFullDocumentDiagnosticReport) - return json.Unmarshal(data, o.FullDocumentDiagnosticReport) - case `"unchanged"`: - o.UnchangedDocumentDiagnosticReport = new(RelatedUnchangedDocumentDiagnosticReport) - return json.Unmarshal(data, o.UnchangedDocumentDiagnosticReport) + switch target := target.(type) { + case *RelatedFullDocumentDiagnosticReport: + o.FullDocumentDiagnosticReport = target + case *RelatedUnchangedDocumentDiagnosticReport: + o.UnchangedDocumentDiagnosticReport = target } - return errInvalidValue("RelatedFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport", data) + return nil } type InlineCompletionListOrItemsOrNull struct { @@ -14701,22 +14746,30 @@ func (o *RequestFailureTelemetryEventOrPerformanceStatsTelemetryEventOrProjectIn _, err := dec.ReadToken() return err case '{': - data, err := dec.ReadValue() + target, err := unmarshalDiscriminatedStruct(dec, "RequestFailureTelemetryEventOrPerformanceStatsTelemetryEventOrProjectInfoTelemetryEventOrNull", "eventName", func(value json.Value) any { + switch string(value) { + case `"languageServer.projectInfo"`: + return new(ProjectInfoTelemetryEvent) + case `"languageServer.errorResponse"`: + return new(RequestFailureTelemetryEvent) + case `"languageServer.performanceStats"`: + return new(PerformanceStatsTelemetryEvent) + default: + return nil + } + }) if err != nil { return err } - switch string(jsonObjectRawField(data, "eventName")) { - case `"languageServer.projectInfo"`: - o.ProjectInfoTelemetryEvent = new(ProjectInfoTelemetryEvent) - return json.Unmarshal(data, o.ProjectInfoTelemetryEvent) - case `"languageServer.errorResponse"`: - o.RequestFailureTelemetryEvent = new(RequestFailureTelemetryEvent) - return json.Unmarshal(data, o.RequestFailureTelemetryEvent) - case `"languageServer.performanceStats"`: - o.PerformanceStatsTelemetryEvent = new(PerformanceStatsTelemetryEvent) - return json.Unmarshal(data, o.PerformanceStatsTelemetryEvent) + switch target := target.(type) { + case *ProjectInfoTelemetryEvent: + o.ProjectInfoTelemetryEvent = target + case *RequestFailureTelemetryEvent: + o.RequestFailureTelemetryEvent = target + case *PerformanceStatsTelemetryEvent: + o.PerformanceStatsTelemetryEvent = target } - return errInvalidValue("RequestFailureTelemetryEventOrPerformanceStatsTelemetryEventOrProjectInfoTelemetryEventOrNull", data) + return nil default: return errInvalidKind("RequestFailureTelemetryEventOrPerformanceStatsTelemetryEventOrProjectInfoTelemetryEventOrNull", dec.PeekKind()) } diff --git a/tsc/internal/lsp/lsproto/lsp_json_benchmark_test.go b/tsc/internal/lsp/lsproto/lsp_json_benchmark_test.go new file mode 100644 index 0000000000000..1b77648ee8022 --- /dev/null +++ b/tsc/internal/lsp/lsproto/lsp_json_benchmark_test.go @@ -0,0 +1,87 @@ +package lsproto + +import ( + "bytes" + "testing" + + "github.com/microsoft/TypeScript/tsc/internal/json" +) + +type bufferedWorkDoneProgressUnion struct { + begin *WorkDoneProgressBegin + report *WorkDoneProgressReport + end *WorkDoneProgressEnd +} + +func (o *bufferedWorkDoneProgressUnion) UnmarshalJSONFrom(dec *json.Decoder) error { + data, err := dec.ReadValue() + if err != nil { + return err + } + switch string(benchmarkRawField(data, "kind")) { + case `"begin"`: + o.begin = new(WorkDoneProgressBegin) + return json.Unmarshal(data, o.begin) + case `"report"`: + o.report = new(WorkDoneProgressReport) + return json.Unmarshal(data, o.report) + case `"end"`: + o.end = new(WorkDoneProgressEnd) + return json.Unmarshal(data, o.end) + default: + return errInvalidValue("bufferedWorkDoneProgressUnion", data) + } +} + +func benchmarkRawField(data []byte, field string) json.Value { + dec := json.NewDecoder(bytes.NewReader(data)) + if _, err := dec.ReadToken(); err != nil { + return nil + } + for dec.PeekKind() != '}' { + name, err := dec.ReadValue() + if err != nil { + return nil + } + if jsonKeyCheck(name, field) { + value, err := dec.ReadValue() + if err != nil { + return nil + } + return value + } + if err := dec.SkipValue(); err != nil { + return nil + } + } + return nil +} + +func BenchmarkUnmarshalDiscriminatedUnion(b *testing.B) { + inputs := map[string][]byte{ + "discriminator-first": []byte(`{"kind":"begin","title":"Indexing","cancellable":true,"message":"Scanning files","percentage":25}`), + "discriminator-last": []byte(`{"title":"Indexing","cancellable":true,"message":"Scanning files","percentage":25,"kind":"begin"}`), + } + for order, input := range inputs { + b.Run(order, func(b *testing.B) { + b.Run("buffered", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + var value bufferedWorkDoneProgressUnion + if err := json.Unmarshal(input, &value); err != nil { + b.Fatal(err) + } + } + }) + b.Run("streaming", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + var value WorkDoneProgressBeginOrReportOrEnd + if err := json.Unmarshal(input, &value); err != nil { + b.Fatal(err) + } + } + }) + }) + } +} diff --git a/tsc/internal/lsp/lsproto/lsp_json_test.go b/tsc/internal/lsp/lsproto/lsp_json_test.go index 3f88840c0b43e..6d7a559bc8647 100644 --- a/tsc/internal/lsp/lsproto/lsp_json_test.go +++ b/tsc/internal/lsp/lsproto/lsp_json_test.go @@ -630,12 +630,30 @@ func TestUnmarshalDiscriminatorUnion(t *testing.T) { assert.Assert(t, v.End != nil) }) + t.Run("discriminator after variant fields", func(t *testing.T) { + t.Parallel() + var v WorkDoneProgressBeginOrReportOrEnd + err := json.Unmarshal([]byte(`{"title": "Indexing", "percentage": 25, "kind": "begin"}`), &v) + assert.NilError(t, err) + assert.Assert(t, v.Begin != nil) + assert.Equal(t, v.Begin.Title, "Indexing") + assert.Assert(t, v.Begin.Percentage != nil) + assert.Equal(t, *v.Begin.Percentage, uint32(25)) + }) + t.Run("invalid discriminator", func(t *testing.T) { t.Parallel() var v WorkDoneProgressBeginOrReportOrEnd err := json.Unmarshal([]byte(`{"kind": "invalid"}`), &v) assert.Assert(t, err != nil) }) + + t.Run("missing discriminator", func(t *testing.T) { + t.Parallel() + var v WorkDoneProgressBeginOrReportOrEnd + err := json.Unmarshal([]byte(`{"message": "missing kind"}`), &v) + assert.ErrorContains(t, err, `missing discriminator "kind"`) + }) } func TestUnmarshalPresenceDiscriminatorUnion(t *testing.T) { @@ -721,6 +739,15 @@ func TestUnmarshalDocumentEditUnion(t *testing.T) { assert.Equal(t, v.CreateFile.Uri, DocumentUri("file:///new.ts")) }) + t.Run("CreateFile with kind after fields", func(t *testing.T) { + t.Parallel() + var v TextDocumentEditOrCreateFileOrRenameFileOrDeleteFile + err := json.Unmarshal([]byte(`{"uri": "file:///new.ts", "kind": "create"}`), &v) + assert.NilError(t, err) + assert.Assert(t, v.CreateFile != nil) + assert.Equal(t, v.CreateFile.Uri, DocumentUri("file:///new.ts")) + }) + t.Run("RenameFile with kind rename", func(t *testing.T) { t.Parallel() var v TextDocumentEditOrCreateFileOrRenameFileOrDeleteFile diff --git a/tsc/internal/lsp/lsproto/structcodec.go b/tsc/internal/lsp/lsproto/structcodec.go index 269364a0f5a86..e1c655e92c516 100644 --- a/tsc/internal/lsp/lsproto/structcodec.go +++ b/tsc/internal/lsp/lsproto/structcodec.go @@ -1,6 +1,7 @@ package lsproto import ( + "fmt" "reflect" "strings" "sync" @@ -31,6 +32,12 @@ type structSpec struct { requiredMask uint64 } +type structDecoder struct { + rv reflect.Value + spec *structSpec + seen uint64 +} + var structSpecCache sync.Map // reflect.Type -> *structSpec func specFor(t reflect.Type) *structSpec { @@ -68,12 +75,46 @@ func specFor(t reflect.Type) *structSpec { return actual.(*structSpec) } +func newStructDecoder(v any) *structDecoder { + rv := reflect.ValueOf(v).Elem() + return &structDecoder{ + rv: rv, + spec: specFor(rv.Type()), + } +} + +func (d *structDecoder) field(name string, kind json.Kind, decode func(any) error) error { + fs, ok := d.spec.byName[name] + if !ok { + return decode(nil) + } + if fs.requiredID >= 0 { + d.seen |= 1 << fs.requiredID + } + if fs.rejectNull && kind == 'n' { + return errNull(name) + } + return decode(d.rv.Field(fs.index).Addr().Interface()) +} + +func (d *structDecoder) finish() error { + if missing := d.spec.requiredMask &^ d.seen; missing != 0 { + var missingProps []string + for id, n := range d.spec.requiredNames { + if missing&(1<= 0 { - seen |= 1 << fs.requiredID + } + if _, err := dec.ReadToken(); err != nil { + return err + } + + return d.finish() +} + +type deferredStructField struct { + name string + value json.Value +} + +// unmarshalDiscriminatedStruct selects a concrete struct using one of its +// fields. Fields after the discriminator decode directly from dec; only fields +// that precede it are retained and replayed. +func unmarshalDiscriminatedStruct( + dec *json.Decoder, + typeName string, + discriminator string, + newTarget func(json.Value) any, +) (any, error) { + if k := dec.PeekKind(); k != '{' { + return nil, errNotObject(k) + } + if _, err := dec.ReadToken(); err != nil { + return nil, err + } + + var deferred []deferredStructField + var target any + var targetDecoder *structDecoder + for dec.PeekKind() != '}' { + rawName, err := dec.ReadValue() + if err != nil { + return nil, err } - if fs.rejectNull && dec.PeekKind() == 'n' { - return errNull(string(name[1 : len(name)-1])) + name := string(rawName[1 : len(rawName)-1]) + + if targetDecoder == nil { + value, err := dec.ReadValue() + if err != nil { + return nil, err + } + if name != discriminator { + deferred = append(deferred, deferredStructField{name: name, value: value.Clone()}) + continue + } + + target = newTarget(value) + if target == nil { + return nil, fmt.Errorf("invalid %s discriminator %q: %s", typeName, discriminator, value) + } + targetDecoder = newStructDecoder(target) + // newTarget validates the string literal discriminator. Mark it as + // seen without reparsing it into the generated zero-sized literal type. + if err := targetDecoder.field(name, value.Kind(), func(any) error { return nil }); err != nil { + return nil, err + } + for _, field := range deferred { + if err := targetDecoder.field(field.name, field.value.Kind(), func(out any) error { + if out == nil { + return nil + } + return json.Unmarshal(field.value, out) + }); err != nil { + return nil, err + } + } + deferred = nil + continue } - if err := json.UnmarshalDecode(dec, rv.Field(fs.index).Addr().Interface()); err != nil { - return err + + if err := targetDecoder.field(name, dec.PeekKind(), func(out any) error { + if out == nil { + return dec.SkipValue() + } + return json.UnmarshalDecode(dec, out) + }); err != nil { + return nil, err } } if _, err := dec.ReadToken(); err != nil { - return err + return nil, err } - if missing := spec.requiredMask &^ seen; missing != 0 { - var missingProps []string - for id, n := range spec.requiredNames { - if missing&(1< Date: Thu, 20 Aug 2026 19:15:00 -0700 Subject: [PATCH 2/3] Better --- .../lsp/lsproto/_generate/generate.mts | 32 +-- tsc/internal/lsp/lsproto/lsp_generated.go | 194 ++++++------------ tsc/internal/lsp/lsproto/structcodec.go | 176 +++++++++------- 3 files changed, 180 insertions(+), 222 deletions(-) diff --git a/tsc/internal/lsp/lsproto/_generate/generate.mts b/tsc/internal/lsp/lsproto/_generate/generate.mts index fe5bb9230c0b5..16bfcb8c4debc 100755 --- a/tsc/internal/lsp/lsproto/_generate/generate.mts +++ b/tsc/internal/lsp/lsproto/_generate/generate.mts @@ -2206,35 +2206,23 @@ function generateCode() { disc: NonNullable>, indent: string, ) { - writeLine(`${indent}target, err := unmarshalDiscriminatedStruct(dec, "${name}", ${JSON.stringify(disc.fieldName)}, func(value json.Value) any {`); - writeLine(`${indent}\tswitch string(value) {`); + writeLine(`${indent}state, err := scanDiscriminatedStruct(dec, "${name}", ${JSON.stringify(disc.fieldName)})`); + writeLine(`${indent}if err != nil {`); + writeLine(`${indent}\treturn err`); + writeLine(`${indent}}`); + writeLine(`${indent}switch string(state.discriminatorValue) {`); for (const [value, entry] of disc.mapping) { - writeLine(`${indent}\tcase \`"${value}"\`:`); - writeLine(`${indent}\t\treturn new(${entry.typeName})`); + writeLine(`${indent}case \`"${value}"\`:`); + writeLine(`${indent}\treturn unmarshalDiscriminatedArm(state, &o.${entry.fieldName})`); } - writeLine(`${indent}\tdefault:`); + writeLine(`${indent}default:`); if (disc.unmapped.length === 1) { - writeLine(`${indent}\t\treturn new(${disc.unmapped[0].typeName})`); + writeLine(`${indent}\treturn unmarshalDiscriminatedArm(state, &o.${disc.unmapped[0].fieldName})`); } else { - writeLine(`${indent}\t\treturn nil`); - } - writeLine(`${indent}\t}`); - writeLine(`${indent}})`); - writeLine(`${indent}if err != nil {`); - writeLine(`${indent}\treturn err`); - writeLine(`${indent}}`); - writeLine(`${indent}switch target := target.(type) {`); - for (const entry of disc.mapping.values()) { - writeLine(`${indent}case *${entry.typeName}:`); - writeLine(`${indent}\to.${entry.fieldName} = target`); - } - for (const entry of disc.unmapped) { - writeLine(`${indent}case *${entry.typeName}:`); - writeLine(`${indent}\to.${entry.fieldName} = target`); + writeLine(`${indent}\treturn state.invalidDiscriminator()`); } writeLine(`${indent}}`); - writeLine(`${indent}return nil`); } /** diff --git a/tsc/internal/lsp/lsproto/lsp_generated.go b/tsc/internal/lsp/lsproto/lsp_generated.go index 9c631db1aa18c..70e75d7fc00d5 100644 --- a/tsc/internal/lsp/lsproto/lsp_generated.go +++ b/tsc/internal/lsp/lsproto/lsp_generated.go @@ -11891,32 +11891,20 @@ var _ json.UnmarshalerFrom = (*TextDocumentEditOrCreateFileOrRenameFileOrDeleteF func (o *TextDocumentEditOrCreateFileOrRenameFileOrDeleteFile) UnmarshalJSONFrom(dec *json.Decoder) error { *o = TextDocumentEditOrCreateFileOrRenameFileOrDeleteFile{} - target, err := unmarshalDiscriminatedStruct(dec, "TextDocumentEditOrCreateFileOrRenameFileOrDeleteFile", "kind", func(value json.Value) any { - switch string(value) { - case `"rename"`: - return new(RenameFile) - case `"create"`: - return new(CreateFile) - case `"delete"`: - return new(DeleteFile) - default: - return new(TextDocumentEdit) - } - }) + state, err := scanDiscriminatedStruct(dec, "TextDocumentEditOrCreateFileOrRenameFileOrDeleteFile", "kind") if err != nil { return err } - switch target := target.(type) { - case *RenameFile: - o.RenameFile = target - case *CreateFile: - o.CreateFile = target - case *DeleteFile: - o.DeleteFile = target - case *TextDocumentEdit: - o.TextDocumentEdit = target + switch string(state.discriminatorValue) { + case `"rename"`: + return unmarshalDiscriminatedArm(state, &o.RenameFile) + case `"create"`: + return unmarshalDiscriminatedArm(state, &o.CreateFile) + case `"delete"`: + return unmarshalDiscriminatedArm(state, &o.DeleteFile) + default: + return unmarshalDiscriminatedArm(state, &o.TextDocumentEdit) } - return nil } type StringOrInlayHintLabelParts struct { @@ -11991,26 +11979,18 @@ var _ json.UnmarshalerFrom = (*WorkspaceFullDocumentDiagnosticReportOrUnchangedD func (o *WorkspaceFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport) UnmarshalJSONFrom(dec *json.Decoder) error { *o = WorkspaceFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport{} - target, err := unmarshalDiscriminatedStruct(dec, "WorkspaceFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport", "kind", func(value json.Value) any { - switch string(value) { - case `"full"`: - return new(WorkspaceFullDocumentDiagnosticReport) - case `"unchanged"`: - return new(WorkspaceUnchangedDocumentDiagnosticReport) - default: - return nil - } - }) + state, err := scanDiscriminatedStruct(dec, "WorkspaceFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport", "kind") if err != nil { return err } - switch target := target.(type) { - case *WorkspaceFullDocumentDiagnosticReport: - o.FullDocumentDiagnosticReport = target - case *WorkspaceUnchangedDocumentDiagnosticReport: - o.UnchangedDocumentDiagnosticReport = target + switch string(state.discriminatorValue) { + case `"full"`: + return unmarshalDiscriminatedArm(state, &o.FullDocumentDiagnosticReport) + case `"unchanged"`: + return unmarshalDiscriminatedArm(state, &o.UnchangedDocumentDiagnosticReport) + default: + return state.invalidDiscriminator() } - return nil } type StringOrStringValue struct { @@ -12383,30 +12363,20 @@ var _ json.UnmarshalerFrom = (*WorkDoneProgressBeginOrReportOrEnd)(nil) func (o *WorkDoneProgressBeginOrReportOrEnd) UnmarshalJSONFrom(dec *json.Decoder) error { *o = WorkDoneProgressBeginOrReportOrEnd{} - target, err := unmarshalDiscriminatedStruct(dec, "WorkDoneProgressBeginOrReportOrEnd", "kind", func(value json.Value) any { - switch string(value) { - case `"begin"`: - return new(WorkDoneProgressBegin) - case `"report"`: - return new(WorkDoneProgressReport) - case `"end"`: - return new(WorkDoneProgressEnd) - default: - return nil - } - }) + state, err := scanDiscriminatedStruct(dec, "WorkDoneProgressBeginOrReportOrEnd", "kind") if err != nil { return err } - switch target := target.(type) { - case *WorkDoneProgressBegin: - o.Begin = target - case *WorkDoneProgressReport: - o.Report = target - case *WorkDoneProgressEnd: - o.End = target + switch string(state.discriminatorValue) { + case `"begin"`: + return unmarshalDiscriminatedArm(state, &o.Begin) + case `"report"`: + return unmarshalDiscriminatedArm(state, &o.Report) + case `"end"`: + return unmarshalDiscriminatedArm(state, &o.End) + default: + return state.invalidDiscriminator() } - return nil } type TextEditOrAnnotatedTextEditOrSnippetTextEdit struct { @@ -12459,26 +12429,18 @@ var _ json.UnmarshalerFrom = (*FullDocumentDiagnosticReportOrUnchangedDocumentDi func (o *FullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport) UnmarshalJSONFrom(dec *json.Decoder) error { *o = FullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport{} - target, err := unmarshalDiscriminatedStruct(dec, "FullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport", "kind", func(value json.Value) any { - switch string(value) { - case `"full"`: - return new(FullDocumentDiagnosticReport) - case `"unchanged"`: - return new(UnchangedDocumentDiagnosticReport) - default: - return nil - } - }) + state, err := scanDiscriminatedStruct(dec, "FullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport", "kind") if err != nil { return err } - switch target := target.(type) { - case *FullDocumentDiagnosticReport: - o.FullDocumentDiagnosticReport = target - case *UnchangedDocumentDiagnosticReport: - o.UnchangedDocumentDiagnosticReport = target + switch string(state.discriminatorValue) { + case `"full"`: + return unmarshalDiscriminatedArm(state, &o.FullDocumentDiagnosticReport) + case `"unchanged"`: + return unmarshalDiscriminatedArm(state, &o.UnchangedDocumentDiagnosticReport) + default: + return state.invalidDiscriminator() } - return nil } type TextDocumentSyncOptionsOrKind struct { @@ -13620,30 +13582,20 @@ var _ json.UnmarshalerFrom = (*VSImageElementOrClassifiedTextElementOrContainerE func (o *VSImageElementOrClassifiedTextElementOrContainerElement) UnmarshalJSONFrom(dec *json.Decoder) error { *o = VSImageElementOrClassifiedTextElementOrContainerElement{} - target, err := unmarshalDiscriminatedStruct(dec, "VSImageElementOrClassifiedTextElementOrContainerElement", "_vs_type", func(value json.Value) any { - switch string(value) { - case `"ContainerElement"`: - return new(VSContainerElement) - case `"ImageElement"`: - return new(VSImageElement) - case `"ClassifiedTextElement"`: - return new(VSClassifiedTextElement) - default: - return nil - } - }) + state, err := scanDiscriminatedStruct(dec, "VSImageElementOrClassifiedTextElementOrContainerElement", "_vs_type") if err != nil { return err } - switch target := target.(type) { - case *VSContainerElement: - o.ContainerElement = target - case *VSImageElement: - o.ImageElement = target - case *VSClassifiedTextElement: - o.ClassifiedTextElement = target + switch string(state.discriminatorValue) { + case `"ContainerElement"`: + return unmarshalDiscriminatedArm(state, &o.ContainerElement) + case `"ImageElement"`: + return unmarshalDiscriminatedArm(state, &o.ImageElement) + case `"ClassifiedTextElement"`: + return unmarshalDiscriminatedArm(state, &o.ClassifiedTextElement) + default: + return state.invalidDiscriminator() } - return nil } type LocationOrLocationsOrDefinitionLinksOrNull struct { @@ -14123,26 +14075,18 @@ var _ json.UnmarshalerFrom = (*RelatedFullDocumentDiagnosticReportOrUnchangedDoc func (o *RelatedFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport) UnmarshalJSONFrom(dec *json.Decoder) error { *o = RelatedFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport{} - target, err := unmarshalDiscriminatedStruct(dec, "RelatedFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport", "kind", func(value json.Value) any { - switch string(value) { - case `"full"`: - return new(RelatedFullDocumentDiagnosticReport) - case `"unchanged"`: - return new(RelatedUnchangedDocumentDiagnosticReport) - default: - return nil - } - }) + state, err := scanDiscriminatedStruct(dec, "RelatedFullDocumentDiagnosticReportOrUnchangedDocumentDiagnosticReport", "kind") if err != nil { return err } - switch target := target.(type) { - case *RelatedFullDocumentDiagnosticReport: - o.FullDocumentDiagnosticReport = target - case *RelatedUnchangedDocumentDiagnosticReport: - o.UnchangedDocumentDiagnosticReport = target + switch string(state.discriminatorValue) { + case `"full"`: + return unmarshalDiscriminatedArm(state, &o.FullDocumentDiagnosticReport) + case `"unchanged"`: + return unmarshalDiscriminatedArm(state, &o.UnchangedDocumentDiagnosticReport) + default: + return state.invalidDiscriminator() } - return nil } type InlineCompletionListOrItemsOrNull struct { @@ -14746,30 +14690,20 @@ func (o *RequestFailureTelemetryEventOrPerformanceStatsTelemetryEventOrProjectIn _, err := dec.ReadToken() return err case '{': - target, err := unmarshalDiscriminatedStruct(dec, "RequestFailureTelemetryEventOrPerformanceStatsTelemetryEventOrProjectInfoTelemetryEventOrNull", "eventName", func(value json.Value) any { - switch string(value) { - case `"languageServer.projectInfo"`: - return new(ProjectInfoTelemetryEvent) - case `"languageServer.errorResponse"`: - return new(RequestFailureTelemetryEvent) - case `"languageServer.performanceStats"`: - return new(PerformanceStatsTelemetryEvent) - default: - return nil - } - }) + state, err := scanDiscriminatedStruct(dec, "RequestFailureTelemetryEventOrPerformanceStatsTelemetryEventOrProjectInfoTelemetryEventOrNull", "eventName") if err != nil { return err } - switch target := target.(type) { - case *ProjectInfoTelemetryEvent: - o.ProjectInfoTelemetryEvent = target - case *RequestFailureTelemetryEvent: - o.RequestFailureTelemetryEvent = target - case *PerformanceStatsTelemetryEvent: - o.PerformanceStatsTelemetryEvent = target + switch string(state.discriminatorValue) { + case `"languageServer.projectInfo"`: + return unmarshalDiscriminatedArm(state, &o.ProjectInfoTelemetryEvent) + case `"languageServer.errorResponse"`: + return unmarshalDiscriminatedArm(state, &o.RequestFailureTelemetryEvent) + case `"languageServer.performanceStats"`: + return unmarshalDiscriminatedArm(state, &o.PerformanceStatsTelemetryEvent) + default: + return state.invalidDiscriminator() } - return nil default: return errInvalidKind("RequestFailureTelemetryEventOrPerformanceStatsTelemetryEventOrProjectInfoTelemetryEventOrNull", dec.PeekKind()) } diff --git a/tsc/internal/lsp/lsproto/structcodec.go b/tsc/internal/lsp/lsproto/structcodec.go index e1c655e92c516..6105ba4faa053 100644 --- a/tsc/internal/lsp/lsproto/structcodec.go +++ b/tsc/internal/lsp/lsproto/structcodec.go @@ -114,7 +114,8 @@ func (d *structDecoder) finish() error { // enforcing object-kind, required-field, and non-nullable-field strictness as // declared by lsp struct tags. Up to 64 required fields are supported. func unmarshalStruct(v any, dec *json.Decoder) error { - d := newStructDecoder(v) + rv := reflect.ValueOf(v).Elem() + spec := specFor(rv.Type()) if k := dec.PeekKind(); k != '{' { return errNotObject(k) @@ -123,18 +124,27 @@ func unmarshalStruct(v any, dec *json.Decoder) error { return err } + var seen uint64 for dec.PeekKind() != '}' { name, err := dec.ReadValue() if err != nil { return err } // name includes surrounding quotes; m[string(b)] is a no-alloc lookup. - if err := d.field(string(name[1:len(name)-1]), dec.PeekKind(), func(out any) error { - if out == nil { - return dec.SkipValue() + fs, ok := spec.byName[string(name[1:len(name)-1])] + if !ok { + if err := dec.SkipValue(); err != nil { + return err } - return json.UnmarshalDecode(dec, out) - }); err != nil { + continue + } + if fs.requiredID >= 0 { + seen |= 1 << fs.requiredID + } + if fs.rejectNull && dec.PeekKind() == 'n' { + return errNull(string(name[1 : len(name)-1])) + } + if err := json.UnmarshalDecode(dec, rv.Field(fs.index).Addr().Interface()); err != nil { return err } } @@ -142,7 +152,16 @@ func unmarshalStruct(v any, dec *json.Decoder) error { return err } - return d.finish() + if missing := spec.requiredMask &^ seen; missing != 0 { + var missingProps []string + for id, n := range spec.requiredNames { + if missing&(1< Date: Thu, 20 Aug 2026 20:56:54 -0700 Subject: [PATCH 3/3] Handle discriminated union fallbacks safely --- .../lsp/lsproto/_generate/generate.mts | 19 +++++++++++----- tsc/internal/lsp/lsproto/lsp_generated.go | 2 +- tsc/internal/lsp/lsproto/lsp_json_test.go | 22 +++++++++++++++++++ tsc/internal/lsp/lsproto/structcodec.go | 21 +++++++++++++----- 4 files changed, 53 insertions(+), 11 deletions(-) diff --git a/tsc/internal/lsp/lsproto/_generate/generate.mts b/tsc/internal/lsp/lsproto/_generate/generate.mts index 16bfcb8c4debc..4e97c69f1e00c 100755 --- a/tsc/internal/lsp/lsproto/_generate/generate.mts +++ b/tsc/internal/lsp/lsproto/_generate/generate.mts @@ -2217,7 +2217,7 @@ function generateCode() { } writeLine(`${indent}default:`); if (disc.unmapped.length === 1) { - writeLine(`${indent}\treturn unmarshalDiscriminatedArm(state, &o.${disc.unmapped[0].fieldName})`); + writeLine(`${indent}\treturn unmarshalDiscriminatedFallbackArm(state, &o.${disc.unmapped[0].fieldName})`); } else { writeLine(`${indent}\treturn state.invalidDiscriminator()`); @@ -2225,6 +2225,15 @@ function generateCode() { writeLine(`${indent}}`); } + function canStreamDiscriminator( + disc: NonNullable>, + ): boolean { + return disc.unmapped.length <= 1 && disc.unmapped.every(entry => + entry.originalType.kind === "reference" + && model.structures.some(structure => structure.name === entry.originalType.name) + ); + } + /** * Generate try-each fallback code for unmapped entries, chaining into * presence dispatch if possible before falling back to raw try-each. @@ -3412,7 +3421,7 @@ function generateCode() { else { let exhaustive = false; const disc = findDiscriminatorField(entries); - if (disc && disc.unmapped.length <= 1) { + if (disc && canStreamDiscriminator(disc)) { generateStreamingDiscriminatorDispatch(name, disc, "\t\t"); exhaustive = true; } @@ -3424,7 +3433,7 @@ function generateCode() { writeLine(`\t\t\treturn err`); writeLine(`\t\t}`); } - if (disc && disc.unmapped.length > 1) { + if (disc && !canStreamDiscriminator(disc)) { exhaustive = generateDiscriminatorDispatch(disc, "\t\t"); } else if (!disc) { @@ -3457,7 +3466,7 @@ function generateCode() { // unions can still stream; other unions use ReadValue + try-each. let exhaustive = false; const disc = findDiscriminatorField(fieldEntries); - if (disc && disc.unmapped.length <= 1) { + if (disc && canStreamDiscriminator(disc)) { generateStreamingDiscriminatorDispatch(name, disc, "\t"); exhaustive = true; } @@ -3473,7 +3482,7 @@ function generateCode() { writeLine(""); } } - if (disc && disc.unmapped.length > 1) { + if (disc && !canStreamDiscriminator(disc)) { exhaustive = generateDiscriminatorDispatch(disc, "\t"); } else if (!disc) { diff --git a/tsc/internal/lsp/lsproto/lsp_generated.go b/tsc/internal/lsp/lsproto/lsp_generated.go index 70e75d7fc00d5..5d8ede04361a0 100644 --- a/tsc/internal/lsp/lsproto/lsp_generated.go +++ b/tsc/internal/lsp/lsproto/lsp_generated.go @@ -11903,7 +11903,7 @@ func (o *TextDocumentEditOrCreateFileOrRenameFileOrDeleteFile) UnmarshalJSONFrom case `"delete"`: return unmarshalDiscriminatedArm(state, &o.DeleteFile) default: - return unmarshalDiscriminatedArm(state, &o.TextDocumentEdit) + return unmarshalDiscriminatedFallbackArm(state, &o.TextDocumentEdit) } } diff --git a/tsc/internal/lsp/lsproto/lsp_json_test.go b/tsc/internal/lsp/lsproto/lsp_json_test.go index 6d7a559bc8647..d188905342fbf 100644 --- a/tsc/internal/lsp/lsproto/lsp_json_test.go +++ b/tsc/internal/lsp/lsproto/lsp_json_test.go @@ -648,6 +648,13 @@ func TestUnmarshalDiscriminatorUnion(t *testing.T) { assert.Assert(t, err != nil) }) + t.Run("non-string discriminator", func(t *testing.T) { + t.Parallel() + var v WorkDoneProgressBeginOrReportOrEnd + err := json.Unmarshal([]byte(`{"kind": null}`), &v) + assert.Assert(t, err != nil) + }) + t.Run("missing discriminator", func(t *testing.T) { t.Parallel() var v WorkDoneProgressBeginOrReportOrEnd @@ -729,6 +736,21 @@ func TestUnmarshalDocumentEditUnion(t *testing.T) { assert.Assert(t, v.DeleteFile == nil) }) + t.Run("TextDocumentEdit with non-string kind", func(t *testing.T) { + t.Parallel() + var v TextDocumentEditOrCreateFileOrRenameFileOrDeleteFile + err := json.Unmarshal([]byte(`{ + "kind": null, + "textDocument": {"uri": "file:///a.ts", "version": 1}, + "edits": [] + }`), &v) + assert.NilError(t, err) + assert.Assert(t, v.TextDocumentEdit != nil) + assert.Assert(t, v.CreateFile == nil) + assert.Assert(t, v.RenameFile == nil) + assert.Assert(t, v.DeleteFile == nil) + }) + t.Run("CreateFile with kind create", func(t *testing.T) { t.Parallel() var v TextDocumentEditOrCreateFileOrRenameFileOrDeleteFile diff --git a/tsc/internal/lsp/lsproto/structcodec.go b/tsc/internal/lsp/lsproto/structcodec.go index 6105ba4faa053..7ab70cd734bc6 100644 --- a/tsc/internal/lsp/lsproto/structcodec.go +++ b/tsc/internal/lsp/lsproto/structcodec.go @@ -205,9 +205,6 @@ func scanDiscriminatedStruct(dec *json.Decoder, typeName string, discriminator s if err != nil { return discriminatedStructDecoder{}, err } - if value.Kind() != '"' { - return discriminatedStructDecoder{}, fmt.Errorf("invalid %s discriminator %q: got %v", typeName, discriminator, value.Kind()) - } state.discriminatorValue = value state.hasDiscriminator = true return state, nil @@ -236,11 +233,25 @@ func (d discriminatedStructDecoder) invalidDiscriminator() error { // unmarshalDiscriminatedArm decodes the retained and remaining object fields // into a concrete union arm, assigning it only after the full object succeeds. func unmarshalDiscriminatedArm[T any](state discriminatedStructDecoder, out **T) error { + return unmarshalDiscriminatedArmWithOptions(state, out, false) +} + +// unmarshalDiscriminatedFallbackArm decodes a fallback arm, including its raw +// discriminator in case the fallback structure declares that field. +func unmarshalDiscriminatedFallbackArm[T any](state discriminatedStructDecoder, out **T) error { + return unmarshalDiscriminatedArmWithOptions(state, out, true) +} + +func unmarshalDiscriminatedArmWithOptions[T any](state discriminatedStructDecoder, out **T, decodeDiscriminator bool) error { target := new(T) targetDecoder := newStructDecoder(target) if state.hasDiscriminator { - // The generated switch validated the literal discriminator. - if err := targetDecoder.field(state.discriminator, '"', func(any) error { return nil }); err != nil { + if err := targetDecoder.field(state.discriminator, state.discriminatorValue.Kind(), func(out any) error { + if out == nil || !decodeDiscriminator { + return nil + } + return json.Unmarshal(state.discriminatorValue, out) + }); err != nil { return err } }