mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
refactor: make ChatMessagePart a discriminated union in TypeScript (#23168)
The flat ChatMessagePart interface had 20+ optional fields, preventing
TypeScript from narrowing types on switch(part.type). Each consumer
needed runtime validation, type assertions, or defensive ?. chains.
Add `variants` struct tags to ChatMessagePart fields declaring which
union variants include each field. A codegen mutation in apitypings
reads these tags via reflect and generates per-variant sub-interfaces
(ChatTextPart, ChatReasoningPart, etc.) plus a union type alias.
A test validates every field has a variants tag or is explicitly
excluded, and every part type is covered.
Remove dead frontend code: normalizeBlockType, alias case branches
("thinking", "toolcall", "toolresult"), legacy field fallbacks
(line_number, typedBlock.name/id/input/output), and result_delta
handling. Add test coverage for args_delta streaming, provider_executed
skip logic, and source part parsing.
This commit is contained in:
+37
-20
@@ -98,6 +98,19 @@ const (
|
||||
ChatMessagePartTypeFileReference ChatMessagePartType = "file-reference"
|
||||
)
|
||||
|
||||
// AllChatMessagePartTypes returns all known ChatMessagePartType values.
|
||||
func AllChatMessagePartTypes() []ChatMessagePartType {
|
||||
return []ChatMessagePartType{
|
||||
ChatMessagePartTypeText,
|
||||
ChatMessagePartTypeReasoning,
|
||||
ChatMessagePartTypeToolCall,
|
||||
ChatMessagePartTypeToolResult,
|
||||
ChatMessagePartTypeSource,
|
||||
ChatMessagePartTypeFile,
|
||||
ChatMessagePartTypeFileReference,
|
||||
}
|
||||
}
|
||||
|
||||
// ChatMessagePart is a structured chunk of a chat message.
|
||||
//
|
||||
// WARNING: This type is both an API wire type and a database
|
||||
@@ -106,37 +119,41 @@ const (
|
||||
// changes, and omitempty behavior all affect backward-compatible
|
||||
// deserialization of stored rows. Treat changes to this struct
|
||||
// with the same care as a database migration.
|
||||
//
|
||||
// The variants struct tag declares which discriminated-union
|
||||
// variants include each field in the generated TypeScript. Bare
|
||||
// name = required, ? suffix = optional. Fields without a variants
|
||||
// tag are excluded from the generated union. See
|
||||
// scripts/apitypings/main.go for the codegen that reads these.
|
||||
type ChatMessagePart struct {
|
||||
Type ChatMessagePartType `json:"type"`
|
||||
Text string `json:"text,omitempty"`
|
||||
Text string `json:"text,omitempty" variants:"text,reasoning"`
|
||||
Signature string `json:"signature,omitempty"`
|
||||
ToolCallID string `json:"tool_call_id,omitempty"`
|
||||
ToolName string `json:"tool_name,omitempty"`
|
||||
Args json.RawMessage `json:"args,omitempty"`
|
||||
ArgsDelta string `json:"args_delta,omitempty"`
|
||||
Result json.RawMessage `json:"result,omitempty"`
|
||||
ToolCallID string `json:"tool_call_id,omitempty" variants:"tool-call,tool-result"`
|
||||
ToolName string `json:"tool_name,omitempty" variants:"tool-call,tool-result"`
|
||||
Args json.RawMessage `json:"args,omitempty" variants:"tool-call?"`
|
||||
ArgsDelta string `json:"args_delta,omitempty" variants:"tool-call?"`
|
||||
Result json.RawMessage `json:"result,omitempty" variants:"tool-result?"`
|
||||
ResultDelta string `json:"result_delta,omitempty"`
|
||||
IsError bool `json:"is_error,omitempty"`
|
||||
SourceID string `json:"source_id,omitempty"`
|
||||
URL string `json:"url,omitempty"`
|
||||
Title string `json:"title,omitempty"`
|
||||
MediaType string `json:"media_type,omitempty"`
|
||||
Data []byte `json:"data,omitempty"`
|
||||
FileID uuid.NullUUID `json:"file_id,omitempty" format:"uuid"`
|
||||
// The following fields are only set when Type is
|
||||
// ChatInputPartTypeFileReference.
|
||||
FileName string `json:"file_name,omitempty"`
|
||||
StartLine int `json:"start_line,omitempty"`
|
||||
EndLine int `json:"end_line,omitempty"`
|
||||
IsError bool `json:"is_error,omitempty" variants:"tool-result?"`
|
||||
SourceID string `json:"source_id,omitempty" variants:"source?"`
|
||||
URL string `json:"url,omitempty" variants:"source"`
|
||||
Title string `json:"title,omitempty" variants:"source?"`
|
||||
MediaType string `json:"media_type,omitempty" variants:"file"`
|
||||
Data []byte `json:"data,omitempty" variants:"file?"`
|
||||
FileID uuid.NullUUID `json:"file_id,omitempty" format:"uuid" variants:"file?"`
|
||||
FileName string `json:"file_name,omitempty" variants:"file-reference"`
|
||||
StartLine int `json:"start_line,omitempty" variants:"file-reference"`
|
||||
EndLine int `json:"end_line,omitempty" variants:"file-reference"`
|
||||
// The code content from the diff that was commented on.
|
||||
Content string `json:"content,omitempty"`
|
||||
Content string `json:"content,omitempty" variants:"file-reference"`
|
||||
// ProviderMetadata holds provider-specific response metadata
|
||||
// (e.g. Anthropic cache control hints) as raw JSON. Internal
|
||||
// only: stripped by db2sdk before API responses.
|
||||
ProviderMetadata json.RawMessage `json:"provider_metadata,omitempty" typescript:"-"`
|
||||
// ProviderExecuted indicates the tool call was executed by
|
||||
// the provider (e.g. Anthropic computer use).
|
||||
ProviderExecuted bool `json:"provider_executed,omitempty"`
|
||||
ProviderExecuted bool `json:"provider_executed,omitempty" variants:"tool-call?,tool-result?"`
|
||||
}
|
||||
|
||||
// StripInternal removes internal-only fields that must not be
|
||||
|
||||
@@ -6,6 +6,8 @@ import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -191,6 +193,85 @@ func TestChatMessagePart_StripInternal(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
// TestChatMessagePartVariantTags validates the `variants` struct tags
|
||||
// on ChatMessagePart fields. Every field must either declare variant
|
||||
// membership or be explicitly excluded, and every known part type
|
||||
// must appear in at least one tag.
|
||||
//
|
||||
// If this test fails, edit the variants struct tags on ChatMessagePart
|
||||
// in codersdk/chats.go.
|
||||
func TestChatMessagePartVariantTags(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const editHint = "edit the variants struct tags on ChatMessagePart in codersdk/chats.go"
|
||||
|
||||
// Fields intentionally excluded from all generated variants.
|
||||
// If you add a new field to ChatMessagePart, either add a
|
||||
// variants tag or add it here with a comment explaining why.
|
||||
excludedFields := map[string]string{
|
||||
"type": "discriminant, added automatically by codegen",
|
||||
"signature": "added in #22290, never populated by any code path",
|
||||
"result_delta": "added in #22290, never populated by any code path",
|
||||
"provider_metadata": "internal only, stripped by db2sdk before API responses",
|
||||
}
|
||||
|
||||
knownTypes := make(map[codersdk.ChatMessagePartType]bool)
|
||||
for _, pt := range codersdk.AllChatMessagePartTypes() {
|
||||
knownTypes[pt] = true
|
||||
}
|
||||
|
||||
// Parse all variants tags from the struct and validate them.
|
||||
typ := reflect.TypeOf(codersdk.ChatMessagePart{})
|
||||
coveredTypes := make(map[codersdk.ChatMessagePartType]bool)
|
||||
hasRequired := make(map[codersdk.ChatMessagePartType]bool)
|
||||
|
||||
for i := range typ.NumField() {
|
||||
f := typ.Field(i)
|
||||
jsonTag := f.Tag.Get("json")
|
||||
if jsonTag == "" || jsonTag == "-" {
|
||||
continue
|
||||
}
|
||||
jsonName, _, _ := strings.Cut(jsonTag, ",")
|
||||
|
||||
varTag := f.Tag.Get("variants")
|
||||
if varTag == "" {
|
||||
assert.Contains(t, excludedFields, jsonName,
|
||||
"field %s (json:%q) has no variants tag and is not in excludedFields; %s",
|
||||
f.Name, jsonName, editHint)
|
||||
continue
|
||||
}
|
||||
|
||||
assert.NotEqual(t, "type", jsonName,
|
||||
"the discriminant field must not have a variants tag; %s", editHint)
|
||||
|
||||
for _, entry := range strings.Split(varTag, ",") {
|
||||
isOptional := strings.HasSuffix(entry, "?")
|
||||
typeLit := codersdk.ChatMessagePartType(strings.TrimSuffix(entry, "?"))
|
||||
|
||||
assert.True(t, knownTypes[typeLit],
|
||||
"field %s variants tag references unknown type %q; %s",
|
||||
f.Name, typeLit, editHint)
|
||||
|
||||
coveredTypes[typeLit] = true
|
||||
if !isOptional {
|
||||
hasRequired[typeLit] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Every known type must appear in at least one variants tag.
|
||||
for pt := range knownTypes {
|
||||
assert.True(t, coveredTypes[pt],
|
||||
"ChatMessagePartType %q is not referenced by any variants tag; %s", pt, editHint)
|
||||
}
|
||||
|
||||
// Every variant must have at least one required field.
|
||||
for pt := range coveredTypes {
|
||||
assert.True(t, hasRequired[pt],
|
||||
"variant %q has no required fields (all have ? suffix); %s", pt, editHint)
|
||||
}
|
||||
}
|
||||
|
||||
func TestModelCostConfig_LegacyNumericJSON(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user