Add feedback-aware correction contracts
This commit is contained in:
@@ -533,6 +533,11 @@ func (client *fakeScenesLLMClient) CompleteStructured(ctx context.Context, req c
|
||||
|
||||
func cloneStructuredCompletionRequest(req contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest {
|
||||
req.Inputs = req.Inputs.Clone()
|
||||
correction, err := contracts.CloneSemanticCorrection(req.Correction)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
req.Correction = correction
|
||||
req.Vars = cloneVars(req.Vars)
|
||||
return req
|
||||
}
|
||||
|
||||
@@ -592,6 +592,11 @@ func (client *fakeCombatTurnsLLMClient) CompleteStructured(_ context.Context, re
|
||||
|
||||
func cloneStructuredCompletionRequest(req contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest {
|
||||
req.Inputs = req.Inputs.Clone()
|
||||
correction, err := contracts.CloneSemanticCorrection(req.Correction)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
req.Correction = correction
|
||||
if len(req.Vars) == 0 {
|
||||
req.Vars = nil
|
||||
return req
|
||||
|
||||
@@ -58,6 +58,11 @@ func newExtractor(t *testing.T, client contracts.StructuredLLMClient, references
|
||||
|
||||
func cloneStructuredCompletionRequest(req contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest {
|
||||
req.Inputs = req.Inputs.Clone()
|
||||
correction, err := contracts.CloneSemanticCorrection(req.Correction)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
req.Correction = correction
|
||||
return req
|
||||
}
|
||||
|
||||
|
||||
@@ -88,5 +88,10 @@ func (client *fakeItemsLLMClient) CompleteStructured(ctx context.Context, req co
|
||||
|
||||
func cloneStructuredCompletionRequest(req contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest {
|
||||
req.Inputs = req.Inputs.Clone()
|
||||
correction, err := contracts.CloneSemanticCorrection(req.Correction)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
req.Correction = correction
|
||||
return req
|
||||
}
|
||||
|
||||
@@ -78,5 +78,10 @@ func (client *fakeOccurrencesLLMClient) CompleteStructured(_ context.Context, re
|
||||
|
||||
func cloneStructuredCompletionRequest(req contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest {
|
||||
req.Inputs = req.Inputs.Clone()
|
||||
correction, err := contracts.CloneSemanticCorrection(req.Correction)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
req.Correction = correction
|
||||
return req
|
||||
}
|
||||
|
||||
@@ -85,5 +85,10 @@ func (client *fakeLocationsLLMClient) CompleteStructured(_ context.Context, req
|
||||
|
||||
func cloneStructuredCompletionRequest(req contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest {
|
||||
req.Inputs = req.Inputs.Clone()
|
||||
correction, err := contracts.CloneSemanticCorrection(req.Correction)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
req.Correction = correction
|
||||
return req
|
||||
}
|
||||
|
||||
@@ -452,5 +452,10 @@ func (client *fakeOccurrencesLLMClient) CompleteStructured(_ context.Context, re
|
||||
|
||||
func cloneStructuredCompletionRequest(req contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest {
|
||||
req.Inputs = req.Inputs.Clone()
|
||||
correction, err := contracts.CloneSemanticCorrection(req.Correction)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
req.Correction = correction
|
||||
return req
|
||||
}
|
||||
|
||||
@@ -65,6 +65,11 @@ func mismatchedSourceInputRequest(req contracts.TypedExtractionRequest) contract
|
||||
|
||||
func cloneStructuredCompletionRequest(req contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest {
|
||||
req.Inputs = req.Inputs.Clone()
|
||||
correction, err := contracts.CloneSemanticCorrection(req.Correction)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
req.Correction = correction
|
||||
if len(req.Vars) == 0 {
|
||||
req.Vars = nil
|
||||
return req
|
||||
|
||||
@@ -61,6 +61,11 @@ func mismatchedSourceInputRequest(req contracts.TypedExtractionRequest) contract
|
||||
|
||||
func cloneStructuredCompletionRequest(req contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest {
|
||||
req.Inputs = req.Inputs.Clone()
|
||||
correction, err := contracts.CloneSemanticCorrection(req.Correction)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
req.Correction = correction
|
||||
if len(req.Vars) == 0 {
|
||||
req.Vars = nil
|
||||
return req
|
||||
|
||||
@@ -163,6 +163,11 @@ func (client *fakeSpellsLLMClient) CompleteStructured(_ context.Context, req con
|
||||
|
||||
func cloneStructuredCompletionRequest(req contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest {
|
||||
req.Inputs = req.Inputs.Clone()
|
||||
correction, err := contracts.CloneSemanticCorrection(req.Correction)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
req.Correction = correction
|
||||
if len(req.Vars) == 0 {
|
||||
req.Vars = nil
|
||||
return req
|
||||
|
||||
@@ -366,7 +366,11 @@ type recordingNormalizerClient struct {
|
||||
}
|
||||
|
||||
func (c *recordingNormalizerClient) CompleteStructured(_ context.Context, request contracts.StructuredCompletionRequest, output any) (contracts.StructuredCompletionResponse, error) {
|
||||
c.requests = append(c.requests, request)
|
||||
snapshot, err := contracts.CloneStructuredCompletionRequest(request)
|
||||
if err != nil {
|
||||
return contracts.StructuredCompletionResponse{}, fmt.Errorf("clone recording request: %w", err)
|
||||
}
|
||||
c.requests = append(c.requests, snapshot)
|
||||
if c.err != nil {
|
||||
return contracts.StructuredCompletionResponse{}, c.err
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package locationregistry
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
|
||||
@@ -17,7 +18,11 @@ type recordingLocationNormalizerClient struct {
|
||||
}
|
||||
|
||||
func (c *recordingLocationNormalizerClient) CompleteStructured(_ context.Context, request contracts.StructuredCompletionRequest, output any) (contracts.StructuredCompletionResponse, error) {
|
||||
c.requests = append(c.requests, request)
|
||||
snapshot, err := contracts.CloneStructuredCompletionRequest(request)
|
||||
if err != nil {
|
||||
return contracts.StructuredCompletionResponse{}, fmt.Errorf("clone recording request: %w", err)
|
||||
}
|
||||
c.requests = append(c.requests, snapshot)
|
||||
if c.err != nil {
|
||||
return contracts.StructuredCompletionResponse{}, c.err
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package npcregistry
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
@@ -148,7 +149,11 @@ type recordingNPCNormalizerClient struct {
|
||||
}
|
||||
|
||||
func (c *recordingNPCNormalizerClient) CompleteStructured(_ context.Context, request contracts.StructuredCompletionRequest, output any) (contracts.StructuredCompletionResponse, error) {
|
||||
c.requests = append(c.requests, request)
|
||||
snapshot, err := contracts.CloneStructuredCompletionRequest(request)
|
||||
if err != nil {
|
||||
return contracts.StructuredCompletionResponse{}, fmt.Errorf("clone recording request: %w", err)
|
||||
}
|
||||
c.requests = append(c.requests, snapshot)
|
||||
if c.err != nil {
|
||||
return contracts.StructuredCompletionResponse{}, c.err
|
||||
}
|
||||
|
||||
@@ -375,8 +375,12 @@ func (client *fakeCombatLLMClient) CompleteStructured(ctx context.Context, req c
|
||||
if err := ctx.Err(); err != nil {
|
||||
return contracts.StructuredCompletionResponse{}, err
|
||||
}
|
||||
snapshot, err := contracts.CloneStructuredCompletionRequest(req)
|
||||
if err != nil {
|
||||
return contracts.StructuredCompletionResponse{}, fmt.Errorf("clone fake request: %w", err)
|
||||
}
|
||||
client.mu.Lock()
|
||||
client.requests = append(client.requests, req)
|
||||
client.requests = append(client.requests, snapshot)
|
||||
index := len(client.requests) - 1
|
||||
client.mu.Unlock()
|
||||
if index >= len(client.responses) {
|
||||
|
||||
@@ -240,7 +240,11 @@ type fakeNPCProductionLLMClient struct {
|
||||
}
|
||||
|
||||
func (client *fakeNPCProductionLLMClient) CompleteStructured(_ context.Context, req contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) {
|
||||
client.requests = append(client.requests, req)
|
||||
snapshot, err := contracts.CloneStructuredCompletionRequest(req)
|
||||
if err != nil {
|
||||
return contracts.StructuredCompletionResponse{}, fmt.Errorf("clone fake request: %w", err)
|
||||
}
|
||||
client.requests = append(client.requests, snapshot)
|
||||
var content []byte
|
||||
switch req.PromptID {
|
||||
case npcregistry.PromptID:
|
||||
|
||||
@@ -71,6 +71,11 @@ func responseSourceRefs(startUnitID int, endUnitID int) []spellSourceRefResponse
|
||||
|
||||
func cloneStructuredCompletionRequest(req contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest {
|
||||
req.Inputs = req.Inputs.Clone()
|
||||
correction, err := contracts.CloneSemanticCorrection(req.Correction)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
req.Correction = correction
|
||||
if len(req.Vars) == 0 {
|
||||
req.Vars = nil
|
||||
return req
|
||||
|
||||
Reference in New Issue
Block a user