Implement combat semantics validator
This commit is contained in:
@@ -183,6 +183,8 @@ This stage is small enough for one implementation prompt.
|
||||
|
||||
## Stage 3: Implement The Typed Validator
|
||||
|
||||
✅ Complete
|
||||
|
||||
### Goal
|
||||
|
||||
Implement package-local validator construction, preconditions, structured
|
||||
|
||||
@@ -0,0 +1,233 @@
|
||||
package combat_semantics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/shared"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/validate/scenedescriptions/shape"
|
||||
)
|
||||
|
||||
const (
|
||||
Key = "extract/dnd/scene-descriptions/combat_semantics"
|
||||
ReasonCodeActiveCombatNotClassified = "scene_active_combat_not_classified"
|
||||
ReasonCodeCombatClassificationUnsupported = "scene_combat_classification_unsupported"
|
||||
combatPolicy = "dnd.scene_descriptions.combat_semantics.v1"
|
||||
maximumExplanationRunes = 512
|
||||
)
|
||||
|
||||
type Options struct{}
|
||||
|
||||
type completionResponse struct {
|
||||
Verdict string `json:"verdict"`
|
||||
Explanation string `json:"explanation"`
|
||||
}
|
||||
|
||||
type Validator struct {
|
||||
llm contracts.StructuredLLMClient
|
||||
promptSHA string
|
||||
responseSchemaSHA string
|
||||
}
|
||||
|
||||
var _ contracts.TypedValidator[dnd.SceneDescriptionList] = (*Validator)(nil)
|
||||
var _ contracts.ManifestMetadataProvider = (*Validator)(nil)
|
||||
var _ pipeline.CheckpointFingerprintProvider = (*Validator)(nil)
|
||||
|
||||
func New(llmClient contracts.StructuredLLMClient, _ Options) (*Validator, error) {
|
||||
if llmClient == nil {
|
||||
return nil, validatorErrorf("LLM client must not be nil")
|
||||
}
|
||||
promptSHA, err := promptAssetMetadata()
|
||||
if err != nil {
|
||||
return nil, validatorErrorf("load prompt metadata: %w", err)
|
||||
}
|
||||
responseSchema, err := loadResponseSchema()
|
||||
if err != nil {
|
||||
return nil, validatorErrorf("load response schema: %w", err)
|
||||
}
|
||||
return &Validator{llm: llmClient, promptSHA: promptSHA, responseSchemaSHA: responseSchema.SHA256}, nil
|
||||
}
|
||||
|
||||
func (v *Validator) Name() string { return Key }
|
||||
|
||||
func (v *Validator) ExecutionClass() contracts.ExecutionClass {
|
||||
return contracts.ExecutionClassLLMBacked
|
||||
}
|
||||
|
||||
func (v *Validator) ManifestMetadata() map[string]any {
|
||||
if v == nil {
|
||||
return nil
|
||||
}
|
||||
return map[string]any{
|
||||
"prompt_id": PromptID,
|
||||
"prompt_version": SchemaVersion,
|
||||
"prompt_sha256": v.promptSHA,
|
||||
"response_schema_key": string(ResponseSchemaKey),
|
||||
"response_schema_id": ResponseSchemaID,
|
||||
"response_schema_name": ResponseSchemaName,
|
||||
"response_schema_version": SchemaVersion,
|
||||
"response_schema_sha256": v.responseSchemaSHA,
|
||||
"combat_policy": combatPolicy,
|
||||
}
|
||||
}
|
||||
|
||||
func (v *Validator) CheckpointFingerprints() []pipeline.CheckpointFingerprint {
|
||||
if v == nil {
|
||||
return nil
|
||||
}
|
||||
return []pipeline.CheckpointFingerprint{
|
||||
{Name: "prompt", Value: v.promptSHA},
|
||||
{Name: "response_schema", Value: v.responseSchemaSHA},
|
||||
{Name: "combat_policy", Value: combatPolicy},
|
||||
}
|
||||
}
|
||||
|
||||
func (v *Validator) Validate(ctx context.Context, req contracts.TypedValidationRequest[dnd.SceneDescriptionList]) (contracts.ValidationResult, error) {
|
||||
if v == nil {
|
||||
return contracts.ValidationResult{}, validatorErrorf("validator must not be nil")
|
||||
}
|
||||
if v.llm == nil {
|
||||
return contracts.ValidationResult{}, validatorErrorf("LLM client must not be nil")
|
||||
}
|
||||
sourceInput, scene, err := validateRequest(ctx, req)
|
||||
if err != nil {
|
||||
return contracts.ValidationResult{}, err
|
||||
}
|
||||
|
||||
var response completionResponse
|
||||
_, err = v.llm.CompleteStructured(ctx, contracts.StructuredCompletionRequest{
|
||||
StageName: Key,
|
||||
PromptID: PromptID,
|
||||
PromptVersion: SchemaVersion,
|
||||
ProfileID: req.LLMProfile,
|
||||
SessionID: req.SessionID,
|
||||
StructuredOutputRepairAttempts: req.StructuredOutputRepairAttempts,
|
||||
Inputs: contracts.LLMInputSet{
|
||||
"transcript": shared.TranscriptPromptMaterial(sourceInput),
|
||||
"proposed_kind": contracts.NewLLMInputMaterial("proposed_kind", "text/plain", []byte(scene.Kind), "", ""),
|
||||
},
|
||||
}, &response)
|
||||
if err != nil {
|
||||
return contracts.ValidationResult{}, validatorErrorf("complete structured output: %w", err)
|
||||
}
|
||||
return interpretResponse(response, scene.Kind)
|
||||
}
|
||||
|
||||
func validateRequest(ctx context.Context, req contracts.TypedValidationRequest[dnd.SceneDescriptionList]) (contracts.LLMInputMaterial, dnd.SceneDescription, error) {
|
||||
if ctx == nil {
|
||||
return contracts.LLMInputMaterial{}, dnd.SceneDescription{}, validatorErrorf("context must not be nil")
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return contracts.LLMInputMaterial{}, dnd.SceneDescription{}, validatorErrorf("context error before validation: %w", err)
|
||||
}
|
||||
if req.Stage != string(pipeline.StageExtract) {
|
||||
return contracts.LLMInputMaterial{}, dnd.SceneDescription{}, validatorErrorf("validation stage must be extract, got %q", req.Stage)
|
||||
}
|
||||
sourceInput, err := shared.PrepareChunkExtraction(ctx, contracts.TypedExtractionRequest{Source: req.Source, Chunk: req.Chunk, SourceInput: req.SourceInput})
|
||||
if err != nil {
|
||||
return contracts.LLMInputMaterial{}, dnd.SceneDescription{}, validatorErrorf("prepare chunk source input: %w", err)
|
||||
}
|
||||
if strings.TrimSpace(req.Chunk.ID) == "" || req.Chunk.ID != strings.TrimSpace(req.Chunk.ID) {
|
||||
return contracts.LLMInputMaterial{}, dnd.SceneDescription{}, validatorErrorf("current chunk ID must be nonblank and trimmed")
|
||||
}
|
||||
if req.Chunk.SourceID != req.Source.ID {
|
||||
return contracts.LLMInputMaterial{}, dnd.SceneDescription{}, validatorErrorf("current chunk source ID %q does not match source %q", req.Chunk.SourceID, req.Source.ID)
|
||||
}
|
||||
if err := source.NewDocumentIndex(req.Source).ValidateRef(req.Chunk.Ref); err != nil {
|
||||
return contracts.LLMInputMaterial{}, dnd.SceneDescription{}, validatorErrorf("validate current chunk source range: %w", err)
|
||||
}
|
||||
if err := shape.ValidateForStage(req.Value, req.Stage); err != nil {
|
||||
return contracts.LLMInputMaterial{}, dnd.SceneDescription{}, validatorErrorf("validate scene description shape: %w", err)
|
||||
}
|
||||
scene := req.Value.Scenes[0]
|
||||
if scene.ID != req.Chunk.ID {
|
||||
return contracts.LLMInputMaterial{}, dnd.SceneDescription{}, validatorErrorf("scene ID must equal current chunk ID")
|
||||
}
|
||||
if scene.SourceRef != req.Chunk.Ref {
|
||||
return contracts.LLMInputMaterial{}, dnd.SceneDescription{}, validatorErrorf("scene source range must equal current chunk range")
|
||||
}
|
||||
return sourceInput, scene, nil
|
||||
}
|
||||
|
||||
func interpretResponse(response completionResponse, proposedKind dnd.SceneKind) (contracts.ValidationResult, error) {
|
||||
explanation, err := validateExplanation(response.Explanation)
|
||||
if err != nil {
|
||||
return contracts.ValidationResult{}, err
|
||||
}
|
||||
switch response.Verdict {
|
||||
case "approved":
|
||||
return contracts.ValidationResult{Approved: true}, nil
|
||||
case "combat_should_be_added":
|
||||
if proposedKind == dnd.SceneKindCombat {
|
||||
return contracts.ValidationResult{}, validatorErrorf("combat_should_be_added verdict is inconsistent with proposed combat kind")
|
||||
}
|
||||
return contracts.ValidationResult{
|
||||
Approved: false,
|
||||
ReasonCode: ReasonCodeActiveCombatNotClassified,
|
||||
Message: "The current chunk contains substantive active combat that is not classified as combat.",
|
||||
CorrectionGuidance: "Return kind: combat for this scene. " + explanation,
|
||||
}, nil
|
||||
case "combat_should_be_removed":
|
||||
if proposedKind != dnd.SceneKindCombat {
|
||||
return contracts.ValidationResult{}, validatorErrorf("combat_should_be_removed verdict is inconsistent with proposed non-combat kind")
|
||||
}
|
||||
return contracts.ValidationResult{
|
||||
Approved: false,
|
||||
ReasonCode: ReasonCodeCombatClassificationUnsupported,
|
||||
Message: "The current chunk does not support a combat classification.",
|
||||
CorrectionGuidance: "Choose the appropriate narrative, recap, or meta kind for this scene. " + explanation,
|
||||
}, nil
|
||||
default:
|
||||
return contracts.ValidationResult{}, validatorErrorf("unsupported combat-semantics verdict %q", response.Verdict)
|
||||
}
|
||||
}
|
||||
|
||||
func validateExplanation(value string) (string, error) {
|
||||
if !utf8.ValidString(value) {
|
||||
return "", validatorErrorf("validator explanation must be valid UTF-8")
|
||||
}
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" {
|
||||
return "", validatorErrorf("validator explanation must not be blank")
|
||||
}
|
||||
if trimmed != value {
|
||||
return "", validatorErrorf("validator explanation must be trimmed")
|
||||
}
|
||||
if utf8.RuneCountInString(value) > maximumExplanationRunes {
|
||||
return "", validatorErrorf("validator explanation exceeds %d Unicode code points", maximumExplanationRunes)
|
||||
}
|
||||
return value, nil
|
||||
}
|
||||
|
||||
func Spec() pipeline.ValidatorSpec {
|
||||
return pipeline.ValidatorSpec{Key: Key, ExecutionClass: contracts.ExecutionClassLLMBacked}
|
||||
}
|
||||
|
||||
func Register(registry *pipeline.ValidatorRegistry) error {
|
||||
return pipeline.RegisterTypedValidatorBuilder(registry, dnd.SceneDescriptionListKind, Spec(), validateOptions, func(request pipeline.BuildRequest) (contracts.TypedValidator[dnd.SceneDescriptionList], error) {
|
||||
options, err := DecodeOptions(request.Options)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return New(request.Dependencies.LLM, options)
|
||||
})
|
||||
}
|
||||
|
||||
func DecodeOptions(options map[string]any) (Options, error) {
|
||||
if err := pipeline.RejectUnknownOptions(options); err != nil {
|
||||
return Options{}, err
|
||||
}
|
||||
return Options{}, nil
|
||||
}
|
||||
|
||||
func validateOptions(options map[string]any) error { _, err := DecodeOptions(options); return err }
|
||||
|
||||
func validatorErrorf(format string, args ...any) error {
|
||||
return fmt.Errorf("scene combat-semantics validator: "+format, args...)
|
||||
}
|
||||
@@ -0,0 +1,231 @@
|
||||
package combat_semantics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd"
|
||||
)
|
||||
|
||||
func TestValidatorInterpretsCombatSemanticsResponses(t *testing.T) {
|
||||
request := validRequest(dnd.SceneKindNarrative)
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
kind dnd.SceneKind
|
||||
response completionResponse
|
||||
approved bool
|
||||
reasonCode string
|
||||
guidance string
|
||||
wantError string
|
||||
}{
|
||||
{name: "approve non-combat", kind: dnd.SceneKindNarrative, response: completionResponse{Verdict: "approved", Explanation: "No active encounter occurs."}, approved: true},
|
||||
{name: "approve combat", kind: dnd.SceneKindCombat, response: completionResponse{Verdict: "approved", Explanation: "Initiative and attacks organize the chunk."}, approved: true},
|
||||
{name: "add combat", kind: dnd.SceneKindNarrative, response: completionResponse{Verdict: "combat_should_be_added", Explanation: "The combatants exchange attacks and damage."}, reasonCode: ReasonCodeActiveCombatNotClassified, guidance: "Return kind: combat"},
|
||||
{name: "remove combat", kind: dnd.SceneKindCombat, response: completionResponse{Verdict: "combat_should_be_removed", Explanation: "The group only plans for a possible fight."}, reasonCode: ReasonCodeCombatClassificationUnsupported, guidance: "narrative, recap, or meta"},
|
||||
{name: "inconsistent added", kind: dnd.SceneKindCombat, response: completionResponse{Verdict: "combat_should_be_added", Explanation: "The combatants exchange attacks."}, wantError: "inconsistent"},
|
||||
{name: "inconsistent removed", kind: dnd.SceneKindNarrative, response: completionResponse{Verdict: "combat_should_be_removed", Explanation: "The group only plans."}, wantError: "inconsistent"},
|
||||
{name: "unknown verdict", kind: dnd.SceneKindNarrative, response: completionResponse{Verdict: "uncertain", Explanation: "The group only plans."}, wantError: "unsupported"},
|
||||
{name: "blank explanation", kind: dnd.SceneKindNarrative, response: completionResponse{Verdict: "approved", Explanation: " "}, wantError: "blank"},
|
||||
{name: "untrimmed explanation", kind: dnd.SceneKindNarrative, response: completionResponse{Verdict: "approved", Explanation: " evidence"}, wantError: "trimmed"},
|
||||
{name: "oversized explanation", kind: dnd.SceneKindNarrative, response: completionResponse{Verdict: "approved", Explanation: strings.Repeat("界", 513)}, wantError: "exceeds"},
|
||||
} {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
request.Value.Scenes[0].Kind = test.kind
|
||||
client := &fakeCombatSemanticsClient{response: test.response}
|
||||
validator := newValidator(t, client)
|
||||
result, err := validator.Validate(context.Background(), request)
|
||||
if test.wantError != "" {
|
||||
if err == nil || !strings.Contains(err.Error(), test.wantError) {
|
||||
t.Fatalf("Validate() error = %v, want %q", err, test.wantError)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil || result.Approved != test.approved || result.ReasonCode != test.reasonCode || !strings.Contains(result.CorrectionGuidance, test.guidance) {
|
||||
t.Fatalf("Validate() = %#v, %v", result, err)
|
||||
}
|
||||
if !result.Approved && (strings.Contains(result.CorrectionGuidance, Key) || strings.Contains(result.CorrectionGuidance, result.ReasonCode) || strings.Contains(result.CorrectionGuidance, request.Chunk.ID)) {
|
||||
t.Fatalf("correction guidance leaked opaque application detail: %q", result.CorrectionGuidance)
|
||||
}
|
||||
if err := contracts.ValidateValidationResult(result); err != nil {
|
||||
t.Fatalf("ValidateValidationResult(%#v) = %v", result, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidatorUsesOnlyChunkTranscriptAndProposedKind(t *testing.T) {
|
||||
request := validRequest(dnd.SceneKindNarrative)
|
||||
repairAttempts := 2
|
||||
request.LLMProfile = "validator-profile"
|
||||
request.SessionID = "validator-session"
|
||||
request.StructuredOutputRepairAttempts = &repairAttempts
|
||||
request.References = contracts.ReferenceSet{Slots: map[string]contracts.ResolvedReferenceSlot{"players": {Items: []contracts.ReferenceItem{{Content: []byte("hidden reference")}}}}}
|
||||
client := &fakeCombatSemanticsClient{response: completionResponse{Verdict: "approved", Explanation: "No active encounter occurs."}}
|
||||
validator := newValidator(t, client)
|
||||
if result, err := validator.Validate(context.Background(), request); err != nil || !result.Approved {
|
||||
t.Fatalf("Validate() = %#v, %v", result, err)
|
||||
}
|
||||
if len(client.requests) != 1 {
|
||||
t.Fatalf("completion count = %d, want 1", len(client.requests))
|
||||
}
|
||||
completed := client.requests[0]
|
||||
if completed.StageName != Key || completed.PromptID != PromptID || completed.PromptVersion != SchemaVersion || completed.ProfileID != request.LLMProfile || completed.SessionID != request.SessionID || completed.StructuredOutputRepairAttempts == request.StructuredOutputRepairAttempts || *completed.StructuredOutputRepairAttempts != repairAttempts {
|
||||
t.Fatalf("completion request = %#v", completed)
|
||||
}
|
||||
if completed.Correction != nil || !reflect.DeepEqual(sortedInputNames(completed.Inputs), []string{"proposed_kind", "transcript"}) || string(completed.Inputs["proposed_kind"].Content) != "narrative" || !reflect.DeepEqual(completed.Inputs["transcript"].Content, request.Chunk.Content) {
|
||||
t.Fatalf("completion inputs = %#v", completed.Inputs)
|
||||
}
|
||||
for _, forbidden := range []string{request.Value.Scenes[0].ID, request.Value.Scenes[0].Title, request.Value.Scenes[0].Summary, "hidden reference"} {
|
||||
for name, input := range completed.Inputs {
|
||||
if strings.Contains(string(input.Content), forbidden) {
|
||||
t.Fatalf("input %q leaked forbidden model-visible value %q", name, forbidden)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidatorRejectsInvalidRequestsAndCompletionFailures(t *testing.T) {
|
||||
base := validRequest(dnd.SceneKindNarrative)
|
||||
canceled, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
ctx context.Context
|
||||
mutate func(*contracts.TypedValidationRequest[dnd.SceneDescriptionList])
|
||||
wantErr string
|
||||
}{
|
||||
{name: "nil context", ctx: nil, wantErr: "context must not be nil"},
|
||||
{name: "canceled context", ctx: canceled, wantErr: "context error"},
|
||||
{name: "wrong stage", ctx: context.Background(), mutate: func(req *contracts.TypedValidationRequest[dnd.SceneDescriptionList]) {
|
||||
req.Stage = string(pipeline.StageNormalize)
|
||||
}, wantErr: "stage must be extract"},
|
||||
{name: "missing source", ctx: context.Background(), mutate: func(req *contracts.TypedValidationRequest[dnd.SceneDescriptionList]) { req.Source = nil }, wantErr: "source must not be nil"},
|
||||
{name: "missing chunk", ctx: context.Background(), mutate: func(req *contracts.TypedValidationRequest[dnd.SceneDescriptionList]) { req.Chunk = nil }, wantErr: "chunk must not be nil"},
|
||||
{name: "mismatched source input", ctx: context.Background(), mutate: func(req *contracts.TypedValidationRequest[dnd.SceneDescriptionList]) {
|
||||
req.SourceInput.Content = []byte(`{"other":true}`)
|
||||
}, wantErr: "must match chunk"},
|
||||
{name: "multiple scenes", ctx: context.Background(), mutate: func(req *contracts.TypedValidationRequest[dnd.SceneDescriptionList]) {
|
||||
req.Value.Scenes = append(req.Value.Scenes, req.Value.Scenes[0])
|
||||
}, wantErr: "exactly one"},
|
||||
{name: "wrong scene ID", ctx: context.Background(), mutate: func(req *contracts.TypedValidationRequest[dnd.SceneDescriptionList]) {
|
||||
req.Value.Scenes[0].ID = "other"
|
||||
}, wantErr: "scene ID"},
|
||||
{name: "wrong scene range", ctx: context.Background(), mutate: func(req *contracts.TypedValidationRequest[dnd.SceneDescriptionList]) {
|
||||
req.Value.Scenes[0].SourceRef.EndUnitID = 1
|
||||
}, wantErr: "scene source range"},
|
||||
} {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
request := cloneValidationRequest(base)
|
||||
if test.mutate != nil {
|
||||
test.mutate(&request)
|
||||
}
|
||||
client := &fakeCombatSemanticsClient{response: completionResponse{Verdict: "approved", Explanation: "No active encounter occurs."}}
|
||||
_, err := newValidator(t, client).Validate(test.ctx, request)
|
||||
if err == nil || !strings.Contains(err.Error(), test.wantErr) || len(client.requests) != 0 {
|
||||
t.Fatalf("Validate() error = %v, calls = %d; want %q and no calls", err, len(client.requests), test.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
client := &fakeCombatSemanticsClient{err: errors.New("provider unavailable")}
|
||||
if _, err := newValidator(t, client).Validate(context.Background(), base); err == nil || !strings.Contains(err.Error(), "complete structured output") {
|
||||
t.Fatalf("completion failure error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidatorConstructionOptionsAndMetadata(t *testing.T) {
|
||||
if _, err := New(nil, Options{}); err == nil || !strings.Contains(err.Error(), "LLM client") {
|
||||
t.Fatalf("New(nil) error = %v", err)
|
||||
}
|
||||
if _, err := DecodeOptions(map[string]any{"unexpected": true}); err == nil {
|
||||
t.Fatal("DecodeOptions() accepted unknown options")
|
||||
}
|
||||
validator := newValidator(t, &fakeCombatSemanticsClient{})
|
||||
if validator.Name() != Key || validator.ExecutionClass() != contracts.ExecutionClassLLMBacked || Spec().CorrectionProtocol != "" {
|
||||
t.Fatalf("validator contract = %#v / %#v", validator, Spec())
|
||||
}
|
||||
metadata, err := json.Marshal(validator.ManifestMetadata())
|
||||
if err != nil || !strings.Contains(string(metadata), ResponseSchemaID) || strings.Contains(string(metadata), "combat_should_be_added") {
|
||||
t.Fatalf("manifest metadata = %s, %v", metadata, err)
|
||||
}
|
||||
fingerprints := validator.CheckpointFingerprints()
|
||||
if len(fingerprints) != 3 || fingerprints[0].Name != "prompt" || fingerprints[1].Name != "response_schema" || fingerprints[2] != (pipeline.CheckpointFingerprint{Name: "combat_policy", Value: combatPolicy}) {
|
||||
t.Fatalf("fingerprints = %#v", fingerprints)
|
||||
}
|
||||
registry := pipeline.NewValidatorRegistry()
|
||||
if err := Register(registry); err != nil {
|
||||
t.Fatalf("Register() error = %v", err)
|
||||
}
|
||||
if spec, ok := registry.Spec(Key); !ok || spec.Key != Key || spec.ExecutionClass != contracts.ExecutionClassLLMBacked {
|
||||
t.Fatalf("registered spec = %#v, present=%t", spec, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func newValidator(t *testing.T, client contracts.StructuredLLMClient) *Validator {
|
||||
t.Helper()
|
||||
validator, err := New(client, Options{})
|
||||
if err != nil {
|
||||
t.Fatalf("New() error = %v", err)
|
||||
}
|
||||
return validator
|
||||
}
|
||||
|
||||
func validRequest(kind dnd.SceneKind) contracts.TypedValidationRequest[dnd.SceneDescriptionList] {
|
||||
document := &source.SourceDocument{ID: "session", Units: []source.SourceUnit{{ID: 1, Text: "The foes attack."}, {ID: 2, Text: "The party responds."}}}
|
||||
chunk := &source.Chunk{ID: "session:chunk:0", SourceID: document.ID, Ref: source.SourceRef{SourceID: document.ID, StartUnitID: 1, EndUnitID: 2}, Content: []byte(`{"units":[1,2]}`), MediaType: "application/json", Units: append([]source.SourceUnit(nil), document.Units...)}
|
||||
return contracts.TypedValidationRequest[dnd.SceneDescriptionList]{Stage: string(pipeline.StageExtract), Source: document, Chunk: chunk, SourceInput: contracts.NewLLMInputMaterial("source", chunk.MediaType, chunk.Content, "sha256:chunk", "file:///session.json"), Value: dnd.SceneDescriptionList{Scenes: []dnd.SceneDescription{{ID: chunk.ID, SourceRef: chunk.Ref, Kind: kind, Title: "A Scene", Summary: "A valid scene."}}}}
|
||||
}
|
||||
|
||||
func cloneValidationRequest(request contracts.TypedValidationRequest[dnd.SceneDescriptionList]) contracts.TypedValidationRequest[dnd.SceneDescriptionList] {
|
||||
request.SourceInput = request.SourceInput.Clone()
|
||||
chunk := *request.Chunk
|
||||
chunk.Content = append([]byte(nil), chunk.Content...)
|
||||
chunk.Units = append([]source.SourceUnit(nil), chunk.Units...)
|
||||
request.Chunk = &chunk
|
||||
request.Value.Scenes = append([]dnd.SceneDescription(nil), request.Value.Scenes...)
|
||||
return request
|
||||
}
|
||||
|
||||
func sortedInputNames(inputs contracts.LLMInputSet) []string {
|
||||
names := make([]string, 0, len(inputs))
|
||||
for name := range inputs {
|
||||
names = append(names, name)
|
||||
}
|
||||
if len(names) == 2 && names[0] == "transcript" {
|
||||
names[0], names[1] = names[1], names[0]
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
type fakeCombatSemanticsClient struct {
|
||||
response completionResponse
|
||||
err error
|
||||
requests []contracts.StructuredCompletionRequest
|
||||
}
|
||||
|
||||
func (client *fakeCombatSemanticsClient) CompleteStructured(_ context.Context, request contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) {
|
||||
cloned, err := contracts.CloneStructuredCompletionRequest(request)
|
||||
if err != nil {
|
||||
return contracts.StructuredCompletionResponse{}, err
|
||||
}
|
||||
client.requests = append(client.requests, cloned)
|
||||
if client.err != nil {
|
||||
return contracts.StructuredCompletionResponse{}, client.err
|
||||
}
|
||||
response, ok := out.(*completionResponse)
|
||||
if !ok {
|
||||
return contracts.StructuredCompletionResponse{}, errors.New("unexpected completion target")
|
||||
}
|
||||
*response = client.response
|
||||
content, err := json.Marshal(client.response)
|
||||
if err != nil {
|
||||
return contracts.StructuredCompletionResponse{}, err
|
||||
}
|
||||
return contracts.StructuredCompletionResponse{Content: content}, nil
|
||||
}
|
||||
Reference in New Issue
Block a user