package integration_test import ( "context" "encoding/json" "fmt" "gitea.maximumdirect.net/eric/notarius/internal/framework/contracts" ) type extractionResponse struct { SpellCasts []spellCastResponse `json:"spell_casts"` } type spellCastResponse struct { Caster string `json:"caster"` Spell string `json:"spell"` SourceRefs []spellSourceRefResponse `json:"source_refs"` } type spellSourceRefResponse struct { StartUnitID int `json:"start_unit_id"` EndUnitID int `json:"end_unit_id"` } type fakeSpellsLLMClient struct { response extractionResponse responses []extractionResponse rawResponses [][]byte requests []contracts.StructuredCompletionRequest } func (client *fakeSpellsLLMClient) CompleteStructured(_ context.Context, req contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) { client.requests = append(client.requests, cloneStructuredCompletionRequest(req)) var content []byte if client.rawResponses != nil { index := len(client.requests) - 1 if index >= len(client.rawResponses) { return contracts.StructuredCompletionResponse{}, fmt.Errorf("missing fake response %d", index) } content = append([]byte(nil), client.rawResponses[index]...) } else { response := client.response if client.responses != nil { index := len(client.requests) - 1 if index >= len(client.responses) { return contracts.StructuredCompletionResponse{}, fmt.Errorf("missing fake response %d", index) } response = client.responses[index] } var err error content, err = json.Marshal(response) if err != nil { return contracts.StructuredCompletionResponse{}, err } } if err := json.Unmarshal(content, out); err != nil { return contracts.StructuredCompletionResponse{}, fmt.Errorf("populate structured target: %w", err) } return contracts.StructuredCompletionResponse{Content: content}, nil } func responseSourceRefs(startUnitID int, endUnitID int) []spellSourceRefResponse { return []spellSourceRefResponse{ { StartUnitID: startUnitID, EndUnitID: endUnitID, }, } } func cloneStructuredCompletionRequest(req contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest { req.Inputs = req.Inputs.Clone() if len(req.Vars) == 0 { req.Vars = nil return req } vars := make(map[string]any, len(req.Vars)) for key, value := range req.Vars { vars[key] = value } req.Vars = vars return req }