Files
notarius/internal/modules/integration/dnd_spells_helpers_test.go

85 lines
2.4 KiB
Go

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
}