@pi-in-go/pigpen-jev 0.1.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/CREDITS.md +22 -0
- package/LICENSE +22 -0
- package/README.md +237 -0
- package/extensions/jev/ask.go +166 -0
- package/extensions/jev/ask_test.go +218 -0
- package/extensions/jev/backend.go +128 -0
- package/extensions/jev/bench_test.go +64 -0
- package/extensions/jev/boundaries_test.go +159 -0
- package/extensions/jev/command.go +224 -0
- package/extensions/jev/commands_test.go +214 -0
- package/extensions/jev/config.go +450 -0
- package/extensions/jev/errors_test.go +191 -0
- package/extensions/jev/extension.go +391 -0
- package/extensions/jev/fakehost_test.go +548 -0
- package/extensions/jev/gate.go +125 -0
- package/extensions/jev/gate_test.go +610 -0
- package/extensions/jev/gatekey_test.go +24 -0
- package/extensions/jev/go.mod +9 -0
- package/extensions/jev/go.sum +2 -0
- package/extensions/jev/go.work +10 -0
- package/extensions/jev/helpers_test.go +404 -0
- package/extensions/jev/memo.go +88 -0
- package/extensions/jev/output.go +89 -0
- package/extensions/jev/output_test.go +187 -0
- package/extensions/jev/ownmodel_test.go +118 -0
- package/extensions/jev/render.go +136 -0
- package/extensions/jev/review_test.go +310 -0
- package/extensions/jev/source_test.go +57 -0
- package/extensions/jev/text.go +174 -0
- package/extensions/jev/trust_test.go +335 -0
- package/extensions/jev/types.go +227 -0
- package/libs/typesafe/CONTRACT.md +125 -0
- package/libs/typesafe/CREDITS.md +37 -0
- package/libs/typesafe/LICENSE +23 -0
- package/libs/typesafe/README.md +19 -0
- package/libs/typesafe/go.mod +3 -0
- package/libs/typesafe/libraries/ownmodel/backend_test.go +496 -0
- package/libs/typesafe/libraries/ownmodel/canon.go +190 -0
- package/libs/typesafe/libraries/ownmodel/convert.go +199 -0
- package/libs/typesafe/libraries/ownmodel/doc.go +15 -0
- package/libs/typesafe/libraries/ownmodel/equivalence_test.go +199 -0
- package/libs/typesafe/libraries/ownmodel/helpers_test.go +155 -0
- package/libs/typesafe/libraries/ownmodel/mutation_test.go +31 -0
- package/libs/typesafe/libraries/ownmodel/ownmodel.go +225 -0
- package/libs/typesafe/libraries/ownmodel/plan.go +442 -0
- package/libs/typesafe/libraries/ownmodel/run.go +288 -0
- package/libs/typesafe/libraries/ownmodel/schema_test.go +254 -0
- package/libs/typesafe/libraries/ownmodel/twins_test.go +169 -0
- package/libs/typesafe/libraries/ownmodel/utils_test.go +125 -0
- package/libs/typesafe/libraries/pigmodel/pigmodel.go +264 -0
- package/libs/typesafe/libraries/pigmodel/pigmodel_test.go +410 -0
- package/libs/typesafe/libraries/typesafe/answers.go +268 -0
- package/libs/typesafe/libraries/typesafe/api_response_test.go +113 -0
- package/libs/typesafe/libraries/typesafe/batch.go +80 -0
- package/libs/typesafe/libraries/typesafe/batch_test.go +133 -0
- package/libs/typesafe/libraries/typesafe/bench_test.go +71 -0
- package/libs/typesafe/libraries/typesafe/client.go +561 -0
- package/libs/typesafe/libraries/typesafe/client_test.go +495 -0
- package/libs/typesafe/libraries/typesafe/crosscheck_test.go +464 -0
- package/libs/typesafe/libraries/typesafe/crosscheck_workflowevals_test.go +219 -0
- package/libs/typesafe/libraries/typesafe/doc.go +27 -0
- package/libs/typesafe/libraries/typesafe/entry.go +142 -0
- package/libs/typesafe/libraries/typesafe/env.go +11 -0
- package/libs/typesafe/libraries/typesafe/errors.go +310 -0
- package/libs/typesafe/libraries/typesafe/errors_test.go +175 -0
- package/libs/typesafe/libraries/typesafe/helpers_test.go +294 -0
- package/libs/typesafe/libraries/typesafe/live_test.go +96 -0
- package/libs/typesafe/libraries/typesafe/logging.go +160 -0
- package/libs/typesafe/libraries/typesafe/logging_test.go +259 -0
- package/libs/typesafe/libraries/typesafe/marshal_test.go +112 -0
- package/libs/typesafe/libraries/typesafe/mutation_test.go +39 -0
- package/libs/typesafe/libraries/typesafe/questions.go +490 -0
- package/libs/typesafe/libraries/typesafe/questions_test.go +166 -0
- package/libs/typesafe/libraries/typesafe/regressions_test.go +159 -0
- package/libs/typesafe/libraries/typesafe/reliability_test.go +649 -0
- package/libs/typesafe/libraries/typesafe/retry.go +350 -0
- package/libs/typesafe/libraries/typesafe/retry_test.go +297 -0
- package/libs/typesafe/libraries/typesafe/runtime_test.go +26 -0
- package/libs/typesafe/libraries/typesafe/transport_test.go +163 -0
- package/libs/typesafe/libraries/typesafe/twins_test.go +127 -0
- package/libs/typesafe/libraries/typesafe/types_test.go +165 -0
- package/libs/typesafe/libraries/typesafe/version.go +10 -0
- package/libs/typesafe/package.json +37 -0
- package/libs/typesafe/provenance.json +49 -0
- package/package.json +42 -0
- package/port/PORT.md +107 -0
- package/port/e2e/gate-and-output.py +35 -0
- package/port/e2e/jev-ask.py +36 -0
- package/port/e2e/model-switch.py +44 -0
- package/port/e2e/off-by-default.py +34 -0
- package/port/gen-scenarios.py +103 -0
- package/port/golden/cache-identical-calls.jsonl +30 -0
- package/port/golden/clear.jsonl +22 -0
- package/port/golden/commands.jsonl +43 -0
- package/port/golden/enforce-accept.jsonl +23 -0
- package/port/golden/enforce-decline.jsonl +22 -0
- package/port/golden/jev-ask.jsonl +20 -0
- package/port/golden/output-advice.jsonl +23 -0
- package/port/golden/output-leak.jsonl +24 -0
- package/port/golden/output-low-confidence.jsonl +22 -0
- package/port/golden/shadow-flagged.jsonl +23 -0
- package/port/golden/unjudged-tools.jsonl +19 -0
- package/port/golden/write-elision.jsonl +21 -0
- package/port/mutate-unit.py +63 -0
- package/port/mutations.json +578 -0
- package/port/oracle/LICENSE +21 -0
- package/port/oracle/README.md +181 -0
- package/port/oracle/SHA256SUMS +8 -0
- package/port/oracle/package.json +43 -0
- package/port/oracle/src/client.ts +409 -0
- package/port/oracle/src/config.ts +363 -0
- package/port/oracle/src/gate.ts +229 -0
- package/port/oracle/src/index.ts +649 -0
- package/port/oracle/src/output.ts +163 -0
- package/port/red-run.log +309 -0
- package/port/scenarios/cache-identical-calls.json +71 -0
- package/port/scenarios/clear.json +61 -0
- package/port/scenarios/commands.json +119 -0
- package/port/scenarios/enforce-accept.json +66 -0
- package/port/scenarios/enforce-decline.json +57 -0
- package/port/scenarios/jev-ask.json +83 -0
- package/port/scenarios/output-advice.json +61 -0
- package/port/scenarios/output-leak.json +61 -0
- package/port/scenarios/output-low-confidence.json +61 -0
- package/port/scenarios/shadow-flagged.json +61 -0
- package/port/scenarios/unjudged-tools.json +55 -0
- package/port/scenarios/write-elision.json +53 -0
- package/provenance.json +18 -0
|
@@ -0,0 +1,410 @@
|
|
|
1
|
+
package pigmodel
|
|
2
|
+
|
|
3
|
+
import (
|
|
4
|
+
"context"
|
|
5
|
+
"errors"
|
|
6
|
+
"go/ast"
|
|
7
|
+
"go/parser"
|
|
8
|
+
"go/token"
|
|
9
|
+
"os"
|
|
10
|
+
"path/filepath"
|
|
11
|
+
"reflect"
|
|
12
|
+
"strconv"
|
|
13
|
+
"strings"
|
|
14
|
+
"sync"
|
|
15
|
+
"testing"
|
|
16
|
+
"time"
|
|
17
|
+
|
|
18
|
+
"github.com/MichaelKinsy/pigpen/components/typesafe/libraries/ownmodel"
|
|
19
|
+
"github.com/MichaelKinsy/pigpen/components/typesafe/libraries/typesafe"
|
|
20
|
+
)
|
|
21
|
+
|
|
22
|
+
// fakeRegistry stands in for the SDK's ModelRegistry: what the host would answer.
|
|
23
|
+
type fakeRegistry struct {
|
|
24
|
+
mu sync.Mutex
|
|
25
|
+
models map[string]map[string]any // "provider/id" -> model
|
|
26
|
+
auth map[string]any
|
|
27
|
+
authErr error
|
|
28
|
+
complete func(model, request, options map[string]any) map[string]any
|
|
29
|
+
|
|
30
|
+
finds []string
|
|
31
|
+
authCalls int
|
|
32
|
+
requests []map[string]any
|
|
33
|
+
options []map[string]any
|
|
34
|
+
modelsIn []map[string]any
|
|
35
|
+
}
|
|
36
|
+
|
|
37
|
+
func (r *fakeRegistry) Find(providerID, modelID string) map[string]any {
|
|
38
|
+
r.mu.Lock()
|
|
39
|
+
defer r.mu.Unlock()
|
|
40
|
+
r.finds = append(r.finds, providerID+"/"+modelID)
|
|
41
|
+
return r.models[providerID+"/"+modelID]
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
func (r *fakeRegistry) GetApiKeyAndHeaders(model map[string]any) (map[string]any, error) {
|
|
45
|
+
r.mu.Lock()
|
|
46
|
+
defer r.mu.Unlock()
|
|
47
|
+
r.authCalls++
|
|
48
|
+
return r.auth, r.authErr
|
|
49
|
+
}
|
|
50
|
+
|
|
51
|
+
func (r *fakeRegistry) Complete(model, request, options map[string]any) map[string]any {
|
|
52
|
+
r.mu.Lock()
|
|
53
|
+
r.requests = append(r.requests, request)
|
|
54
|
+
r.options = append(r.options, options)
|
|
55
|
+
r.modelsIn = append(r.modelsIn, model)
|
|
56
|
+
fn := r.complete
|
|
57
|
+
r.mu.Unlock()
|
|
58
|
+
return fn(model, request, options)
|
|
59
|
+
}
|
|
60
|
+
|
|
61
|
+
func message(text string, usageIn, usageOut float64) map[string]any {
|
|
62
|
+
return map[string]any{
|
|
63
|
+
"role": "assistant", "content": []any{map[string]any{"type": "text", "text": text}},
|
|
64
|
+
"usage": map[string]any{"input": usageIn, "output": usageOut},
|
|
65
|
+
"stopReason": "stop",
|
|
66
|
+
}
|
|
67
|
+
}
|
|
68
|
+
|
|
69
|
+
func newRegistry(reply func(model, request, options map[string]any) map[string]any) *fakeRegistry {
|
|
70
|
+
return &fakeRegistry{
|
|
71
|
+
models: map[string]map[string]any{"prov/mod": {"id": "mod", "provider": "prov", "api": "some-api"}},
|
|
72
|
+
auth: map[string]any{"ok": true, "apiKey": "key-1", "headers": map[string]any{"X-A": "b"}},
|
|
73
|
+
complete: reply,
|
|
74
|
+
}
|
|
75
|
+
}
|
|
76
|
+
|
|
77
|
+
func ref() Ref { return Ref{Provider: "prov", ID: "mod"} }
|
|
78
|
+
|
|
79
|
+
func mustModel(t *testing.T, r Registry) *Model {
|
|
80
|
+
t.Helper()
|
|
81
|
+
m, err := New(r, ref())
|
|
82
|
+
if err != nil {
|
|
83
|
+
t.Fatal(err)
|
|
84
|
+
}
|
|
85
|
+
return m
|
|
86
|
+
}
|
|
87
|
+
|
|
88
|
+
func asErr[T any](t *testing.T, err error) T {
|
|
89
|
+
t.Helper()
|
|
90
|
+
var target T
|
|
91
|
+
if !errors.As(err, &target) {
|
|
92
|
+
t.Fatalf("error %T (%v) is not %T", err, err, target)
|
|
93
|
+
}
|
|
94
|
+
return target
|
|
95
|
+
}
|
|
96
|
+
|
|
97
|
+
func conversation() ownmodel.Request {
|
|
98
|
+
return ownmodel.Request{Messages: []ownmodel.Message{
|
|
99
|
+
{Role: ownmodel.RoleSystem, Content: "be exact"},
|
|
100
|
+
{Role: ownmodel.RoleUser, Content: "<document>x</document>"},
|
|
101
|
+
}, Schema: map[string]any{"type": "object"}}
|
|
102
|
+
}
|
|
103
|
+
|
|
104
|
+
func TestComplete_SendsTheConversationThroughTheRegistry(t *testing.T) {
|
|
105
|
+
r := newRegistry(func(model, request, options map[string]any) map[string]any { return message(`{"answers":{}}`, 12, 5) })
|
|
106
|
+
m := mustModel(t, r)
|
|
107
|
+
res, err := m.Complete(context.Background(), conversation())
|
|
108
|
+
if err != nil {
|
|
109
|
+
t.Fatal(err)
|
|
110
|
+
}
|
|
111
|
+
if res.Text != `{"answers":{}}` || res.InputTokens == nil || *res.InputTokens != 12 || res.OutputTokens == nil || *res.OutputTokens != 5 {
|
|
112
|
+
t.Fatalf("result %+v", res)
|
|
113
|
+
}
|
|
114
|
+
if m.Name() != "mod" {
|
|
115
|
+
t.Fatalf("name %q", m.Name())
|
|
116
|
+
}
|
|
117
|
+
req := r.requests[0]
|
|
118
|
+
if req["systemPrompt"] != "be exact" {
|
|
119
|
+
t.Fatalf("system prompt %v", req["systemPrompt"])
|
|
120
|
+
}
|
|
121
|
+
msgs := req["messages"].([]any)
|
|
122
|
+
if len(msgs) != 1 {
|
|
123
|
+
t.Fatalf("messages %v", msgs)
|
|
124
|
+
}
|
|
125
|
+
um := msgs[0].(map[string]any)
|
|
126
|
+
if um["role"] != "user" {
|
|
127
|
+
t.Fatalf("role %v", um["role"])
|
|
128
|
+
}
|
|
129
|
+
content := um["content"].([]any)[0].(map[string]any)
|
|
130
|
+
if content["type"] != "text" || content["text"] != "<document>x</document>" {
|
|
131
|
+
t.Fatalf("content %v", content)
|
|
132
|
+
}
|
|
133
|
+
if _, ok := um["timestamp"]; !ok {
|
|
134
|
+
t.Fatal("a Pi user message carries a timestamp")
|
|
135
|
+
}
|
|
136
|
+
if r.options[0]["apiKey"] != "key-1" || !reflect.DeepEqual(r.options[0]["headers"], map[string]any{"X-A": "b"}) {
|
|
137
|
+
t.Fatalf("options %v", r.options[0])
|
|
138
|
+
}
|
|
139
|
+
if !reflect.DeepEqual(r.modelsIn[0], r.models["prov/mod"]) {
|
|
140
|
+
t.Fatal("the model handed to Complete must be the registry's")
|
|
141
|
+
}
|
|
142
|
+
}
|
|
143
|
+
|
|
144
|
+
func TestComplete_ResolvesTheModelAndAuthOnFirstUseAndCachesTheModel(t *testing.T) {
|
|
145
|
+
r := newRegistry(func(model, request, options map[string]any) map[string]any { return message("x", 1, 1) })
|
|
146
|
+
m := mustModel(t, r)
|
|
147
|
+
if len(r.finds) != 0 || r.authCalls != 0 {
|
|
148
|
+
t.Fatal("New must not touch the registry")
|
|
149
|
+
}
|
|
150
|
+
for i := 0; i < 2; i++ {
|
|
151
|
+
if _, err := m.Complete(context.Background(), conversation()); err != nil {
|
|
152
|
+
t.Fatal(err)
|
|
153
|
+
}
|
|
154
|
+
}
|
|
155
|
+
if !reflect.DeepEqual(r.finds, []string{"prov/mod"}) {
|
|
156
|
+
t.Fatalf("finds %v", r.finds)
|
|
157
|
+
}
|
|
158
|
+
if r.authCalls != 2 { // credentials are resolved per request: they can refresh
|
|
159
|
+
t.Fatalf("auth calls %d", r.authCalls)
|
|
160
|
+
}
|
|
161
|
+
}
|
|
162
|
+
|
|
163
|
+
func TestComplete_ReportsAModelTheRegistryDoesNotKnow(t *testing.T) {
|
|
164
|
+
r := newRegistry(func(model, request, options map[string]any) map[string]any { return message("x", 1, 1) })
|
|
165
|
+
m, err := New(r, Ref{Provider: "prov", ID: "missing"})
|
|
166
|
+
if err != nil {
|
|
167
|
+
t.Fatal(err)
|
|
168
|
+
}
|
|
169
|
+
_, err = m.Complete(context.Background(), conversation())
|
|
170
|
+
asErr[*typesafe.TypeSafeError](t, err)
|
|
171
|
+
if !strings.Contains(err.Error(), "prov/missing") {
|
|
172
|
+
t.Fatalf("got %v", err)
|
|
173
|
+
}
|
|
174
|
+
}
|
|
175
|
+
|
|
176
|
+
func TestComplete_AuthFailuresAreReported(t *testing.T) {
|
|
177
|
+
r := newRegistry(func(model, request, options map[string]any) map[string]any { return message("x", 1, 1) })
|
|
178
|
+
r.auth = map[string]any{"ok": false, "error": `No API key found for "prov"`}
|
|
179
|
+
_, err := mustModel(t, r).Complete(context.Background(), conversation())
|
|
180
|
+
asErr[*typesafe.TypeSafeError](t, err)
|
|
181
|
+
if !strings.Contains(err.Error(), `No API key found for "prov"`) {
|
|
182
|
+
t.Fatalf("got %v", err)
|
|
183
|
+
}
|
|
184
|
+
r = newRegistry(func(model, request, options map[string]any) map[string]any { return message("x", 1, 1) })
|
|
185
|
+
r.authErr = errors.New("host down")
|
|
186
|
+
_, err = mustModel(t, r).Complete(context.Background(), conversation())
|
|
187
|
+
if err == nil || !strings.Contains(err.Error(), "host down") {
|
|
188
|
+
t.Fatalf("got %v", err)
|
|
189
|
+
}
|
|
190
|
+
}
|
|
191
|
+
|
|
192
|
+
func TestComplete_AssistantTurnsCarryTheModelIdentity(t *testing.T) {
|
|
193
|
+
r := newRegistry(func(model, request, options map[string]any) map[string]any { return message("x", 1, 1) })
|
|
194
|
+
req := conversation()
|
|
195
|
+
req.Messages = append(req.Messages,
|
|
196
|
+
ownmodel.Message{Role: ownmodel.RoleAssistant, Content: `{"answers":`},
|
|
197
|
+
ownmodel.Message{Role: ownmodel.RoleUser, Content: "fix it"})
|
|
198
|
+
if _, err := mustModel(t, r).Complete(context.Background(), req); err != nil {
|
|
199
|
+
t.Fatal(err)
|
|
200
|
+
}
|
|
201
|
+
msgs := r.requests[0]["messages"].([]any)
|
|
202
|
+
if len(msgs) != 3 {
|
|
203
|
+
t.Fatalf("messages %v", msgs)
|
|
204
|
+
}
|
|
205
|
+
a := msgs[1].(map[string]any)
|
|
206
|
+
if a["role"] != "assistant" || a["api"] != "some-api" || a["provider"] != "prov" || a["model"] != "mod" || a["stopReason"] != "stop" {
|
|
207
|
+
t.Fatalf("assistant message %v", a)
|
|
208
|
+
}
|
|
209
|
+
if txt := a["content"].([]any)[0].(map[string]any); txt["type"] != "text" || txt["text"] != `{"answers":` {
|
|
210
|
+
t.Fatalf("assistant content %v", txt)
|
|
211
|
+
}
|
|
212
|
+
if _, ok := a["usage"].(map[string]any); !ok {
|
|
213
|
+
t.Fatal("an assistant message carries usage")
|
|
214
|
+
}
|
|
215
|
+
}
|
|
216
|
+
|
|
217
|
+
func TestComplete_MultipleSystemMessagesAreJoined(t *testing.T) {
|
|
218
|
+
r := newRegistry(func(model, request, options map[string]any) map[string]any { return message("x", 1, 1) })
|
|
219
|
+
req := ownmodel.Request{Messages: []ownmodel.Message{{Role: ownmodel.RoleSystem, Content: "a"}, {Role: ownmodel.RoleSystem, Content: "b"}, {Role: ownmodel.RoleUser, Content: "u"}}}
|
|
220
|
+
if _, err := mustModel(t, r).Complete(context.Background(), req); err != nil {
|
|
221
|
+
t.Fatal(err)
|
|
222
|
+
}
|
|
223
|
+
if r.requests[0]["systemPrompt"] != "a\n\nb" {
|
|
224
|
+
t.Fatalf("system prompt %v", r.requests[0]["systemPrompt"])
|
|
225
|
+
}
|
|
226
|
+
}
|
|
227
|
+
|
|
228
|
+
func TestComplete_TextBlocksAreConcatenatedAndOtherBlocksIgnored(t *testing.T) {
|
|
229
|
+
r := newRegistry(func(model, request, options map[string]any) map[string]any {
|
|
230
|
+
return map[string]any{"role": "assistant", "stopReason": "stop", "content": []any{
|
|
231
|
+
map[string]any{"type": "thinking", "thinking": "hmm"},
|
|
232
|
+
map[string]any{"type": "text", "text": `{"a":`},
|
|
233
|
+
map[string]any{"type": "text", "text": `1}`},
|
|
234
|
+
}}
|
|
235
|
+
})
|
|
236
|
+
res, err := mustModel(t, r).Complete(context.Background(), conversation())
|
|
237
|
+
if err != nil {
|
|
238
|
+
t.Fatal(err)
|
|
239
|
+
}
|
|
240
|
+
if res.Text != `{"a":1}` || res.InputTokens != nil || res.OutputTokens != nil {
|
|
241
|
+
t.Fatalf("result %+v (unreported usage must stay unknown)", res)
|
|
242
|
+
}
|
|
243
|
+
}
|
|
244
|
+
|
|
245
|
+
func TestComplete_StopReasonsError(t *testing.T) {
|
|
246
|
+
cases := []struct {
|
|
247
|
+
name string
|
|
248
|
+
reply map[string]any
|
|
249
|
+
check func(t *testing.T, err error)
|
|
250
|
+
retrys bool
|
|
251
|
+
}{
|
|
252
|
+
{"status in the message", map[string]any{"stopReason": "error", "errorMessage": "429 rate limited"}, func(t *testing.T, err error) {
|
|
253
|
+
asErr[*typesafe.RateLimitError](t, err)
|
|
254
|
+
}, true},
|
|
255
|
+
{"server status", map[string]any{"stopReason": "error", "errorMessage": "503 upstream"}, func(t *testing.T, err error) {
|
|
256
|
+
asErr[*typesafe.InternalServerError](t, err)
|
|
257
|
+
}, true},
|
|
258
|
+
{"client status", map[string]any{"stopReason": "error", "errorMessage": "401 bad key"}, func(t *testing.T, err error) {
|
|
259
|
+
asErr[*typesafe.AuthenticationError](t, err)
|
|
260
|
+
}, false},
|
|
261
|
+
{"no status", map[string]any{"stopReason": "error", "errorMessage": "socket hang up"}, func(t *testing.T, err error) {
|
|
262
|
+
asErr[*typesafe.APIConnectionError](t, err)
|
|
263
|
+
if !strings.Contains(err.Error(), "socket hang up") {
|
|
264
|
+
t.Fatalf("got %v", err)
|
|
265
|
+
}
|
|
266
|
+
}, true},
|
|
267
|
+
{"no message", map[string]any{"stopReason": "error"}, func(t *testing.T, err error) { asErr[*typesafe.APIConnectionError](t, err) }, true},
|
|
268
|
+
}
|
|
269
|
+
for _, tc := range cases {
|
|
270
|
+
r := newRegistry(func(model, request, options map[string]any) map[string]any { return tc.reply })
|
|
271
|
+
_, err := mustModel(t, r).Complete(context.Background(), conversation())
|
|
272
|
+
if err == nil {
|
|
273
|
+
t.Fatalf("%s: want an error", tc.name)
|
|
274
|
+
}
|
|
275
|
+
tc.check(t, err)
|
|
276
|
+
policy := typesafe.DefaultRetryPolicy()
|
|
277
|
+
calls := 0
|
|
278
|
+
_, _ = typesafe.Retry(context.Background(), func() typesafe.RetryPolicy { policy.BackoffInitial = 0; return policy }(), typesafe.RetryHooks{}, func(ctx context.Context, attempt int) (int, error) {
|
|
279
|
+
calls++
|
|
280
|
+
_, err := mustModel(t, r).Complete(ctx, conversation())
|
|
281
|
+
return 0, err
|
|
282
|
+
})
|
|
283
|
+
if (calls > 1) != tc.retrys {
|
|
284
|
+
t.Fatalf("%s: retried=%v, want %v", tc.name, calls > 1, tc.retrys)
|
|
285
|
+
}
|
|
286
|
+
}
|
|
287
|
+
}
|
|
288
|
+
|
|
289
|
+
func TestComplete_StopReasonsNotAnswers(t *testing.T) {
|
|
290
|
+
for reason, want := range map[string]string{
|
|
291
|
+
"length": "output",
|
|
292
|
+
"toolUse": "tool",
|
|
293
|
+
"weird": `"weird"`,
|
|
294
|
+
} {
|
|
295
|
+
r := newRegistry(func(model, request, options map[string]any) map[string]any {
|
|
296
|
+
return map[string]any{"stopReason": reason, "content": []any{map[string]any{"type": "text", "text": "partial"}}}
|
|
297
|
+
})
|
|
298
|
+
_, err := mustModel(t, r).Complete(context.Background(), conversation())
|
|
299
|
+
te := asErr[*typesafe.TypeSafeError](t, err)
|
|
300
|
+
if !strings.Contains(te.Message, want) {
|
|
301
|
+
t.Fatalf("%s: %q lacks %q", reason, te.Message, want)
|
|
302
|
+
}
|
|
303
|
+
var api *typesafe.APIError
|
|
304
|
+
var conn *typesafe.APIConnectionError
|
|
305
|
+
if errors.As(err, &api) || errors.As(err, &conn) {
|
|
306
|
+
t.Fatalf("%s: must not be retryable", reason)
|
|
307
|
+
}
|
|
308
|
+
}
|
|
309
|
+
r := newRegistry(func(model, request, options map[string]any) map[string]any {
|
|
310
|
+
return map[string]any{"stopReason": "aborted"}
|
|
311
|
+
})
|
|
312
|
+
_, err := mustModel(t, r).Complete(context.Background(), conversation())
|
|
313
|
+
asErr[*typesafe.APIUserAbortError](t, err)
|
|
314
|
+
r = newRegistry(func(model, request, options map[string]any) map[string]any { return nil })
|
|
315
|
+
_, err = mustModel(t, r).Complete(context.Background(), conversation())
|
|
316
|
+
asErr[*typesafe.TypeSafeError](t, err)
|
|
317
|
+
}
|
|
318
|
+
|
|
319
|
+
func TestComplete_AContextThatEndsReturnsAtOnceWithoutWaitingForTheHost(t *testing.T) {
|
|
320
|
+
release := make(chan struct{})
|
|
321
|
+
r := newRegistry(func(model, request, options map[string]any) map[string]any { <-release; return message("x", 1, 1) })
|
|
322
|
+
defer close(release)
|
|
323
|
+
ctx, cancel := context.WithCancel(context.Background())
|
|
324
|
+
done := make(chan error, 1)
|
|
325
|
+
go func() { _, err := mustModel(t, r).Complete(ctx, conversation()); done <- err }()
|
|
326
|
+
time.Sleep(20 * time.Millisecond)
|
|
327
|
+
cancel()
|
|
328
|
+
select {
|
|
329
|
+
case err := <-done:
|
|
330
|
+
asErr[*typesafe.APIUserAbortError](t, err)
|
|
331
|
+
if !errors.Is(err, context.Canceled) {
|
|
332
|
+
t.Fatalf("the abort must carry the context error: %v", err)
|
|
333
|
+
}
|
|
334
|
+
case <-time.After(2 * time.Second):
|
|
335
|
+
t.Fatal("Complete did not return after the context ended")
|
|
336
|
+
}
|
|
337
|
+
}
|
|
338
|
+
|
|
339
|
+
func TestComplete_NativeStructuredOutputIsNotSupported(t *testing.T) {
|
|
340
|
+
r := newRegistry(func(model, request, options map[string]any) map[string]any { return message("x", 1, 1) })
|
|
341
|
+
req := conversation()
|
|
342
|
+
req.Structured = true
|
|
343
|
+
_, err := mustModel(t, r).Complete(context.Background(), req)
|
|
344
|
+
asErr[*typesafe.TypeSafeError](t, err)
|
|
345
|
+
if !strings.Contains(err.Error(), "structured") {
|
|
346
|
+
t.Fatalf("got %v", err)
|
|
347
|
+
}
|
|
348
|
+
if len(r.requests) != 0 {
|
|
349
|
+
t.Fatal("no request may be sent")
|
|
350
|
+
}
|
|
351
|
+
}
|
|
352
|
+
|
|
353
|
+
func TestProviderShapesOfTheModelMap(t *testing.T) {
|
|
354
|
+
// The provider may be a string or an object with an id, the model id "id" or "modelId".
|
|
355
|
+
for _, m := range []map[string]any{
|
|
356
|
+
{"id": "mod", "provider": "prov", "api": "x"},
|
|
357
|
+
{"modelId": "mod", "provider": map[string]any{"id": "prov"}, "api": "x"},
|
|
358
|
+
} {
|
|
359
|
+
r := newRegistry(func(model, request, options map[string]any) map[string]any { return message("x", 1, 1) })
|
|
360
|
+
r.models["prov/mod"] = m
|
|
361
|
+
req := conversation()
|
|
362
|
+
req.Messages = append(req.Messages, ownmodel.Message{Role: ownmodel.RoleAssistant, Content: "a"}, ownmodel.Message{Role: ownmodel.RoleUser, Content: "b"})
|
|
363
|
+
if _, err := mustModel(t, r).Complete(context.Background(), req); err != nil {
|
|
364
|
+
t.Fatal(err)
|
|
365
|
+
}
|
|
366
|
+
a := r.requests[0]["messages"].([]any)[1].(map[string]any)
|
|
367
|
+
if a["provider"] != "prov" || a["model"] != "mod" {
|
|
368
|
+
t.Fatalf("model %v -> assistant %v", m, a)
|
|
369
|
+
}
|
|
370
|
+
}
|
|
371
|
+
}
|
|
372
|
+
|
|
373
|
+
func TestNoProviderOrModelNamesAreBuiltIn(t *testing.T) {
|
|
374
|
+
// Provider routing is data: no source string of the own-model packages names a provider
|
|
375
|
+
// or a model family.
|
|
376
|
+
banned := []string{"openai", "anthropic", "gemini", "google", "copilot", "claude", "gpt", "bedrock", "azure", "mistral", "llama", "deepseek"}
|
|
377
|
+
for _, dir := range []string{".", filepath.Join("..", "ownmodel")} {
|
|
378
|
+
files, err := filepath.Glob(filepath.Join(dir, "*.go"))
|
|
379
|
+
if err != nil {
|
|
380
|
+
t.Fatal(err)
|
|
381
|
+
}
|
|
382
|
+
fset := token.NewFileSet()
|
|
383
|
+
for _, file := range files {
|
|
384
|
+
if strings.HasSuffix(file, "_test.go") {
|
|
385
|
+
continue
|
|
386
|
+
}
|
|
387
|
+
src, err := os.ReadFile(file)
|
|
388
|
+
if err != nil {
|
|
389
|
+
t.Fatal(err)
|
|
390
|
+
}
|
|
391
|
+
f, err := parser.ParseFile(fset, file, src, 0) // comments are not source strings
|
|
392
|
+
if err != nil {
|
|
393
|
+
t.Fatal(err)
|
|
394
|
+
}
|
|
395
|
+
ast.Inspect(f, func(n ast.Node) bool {
|
|
396
|
+
lit, ok := n.(*ast.BasicLit)
|
|
397
|
+
if !ok || lit.Kind != token.STRING {
|
|
398
|
+
return true
|
|
399
|
+
}
|
|
400
|
+
s, _ := strconv.Unquote(lit.Value)
|
|
401
|
+
for _, b := range banned {
|
|
402
|
+
if strings.Contains(strings.ToLower(s), b) {
|
|
403
|
+
t.Errorf("%s: string %q names %q", fset.Position(lit.Pos()), s, b)
|
|
404
|
+
}
|
|
405
|
+
}
|
|
406
|
+
return true
|
|
407
|
+
})
|
|
408
|
+
}
|
|
409
|
+
}
|
|
410
|
+
}
|
|
@@ -0,0 +1,268 @@
|
|
|
1
|
+
package typesafe
|
|
2
|
+
|
|
3
|
+
import (
|
|
4
|
+
"encoding/json"
|
|
5
|
+
"fmt"
|
|
6
|
+
"sort"
|
|
7
|
+
)
|
|
8
|
+
|
|
9
|
+
// Answer is the typed answer to one question: [NoulAnswer], [ChoiceAnswer] or
|
|
10
|
+
// [ScoreAnswer]. An answer with an unknown type decodes as [RawAnswer].
|
|
11
|
+
type Answer interface {
|
|
12
|
+
// AnswerType returns the wire type of the answer.
|
|
13
|
+
AnswerType() QuestionType
|
|
14
|
+
isAnswer()
|
|
15
|
+
}
|
|
16
|
+
|
|
17
|
+
// NoulAnswer is a yes/no answer.
|
|
18
|
+
type NoulAnswer struct {
|
|
19
|
+
// Noul is the probability of a yes answer, from zero to one.
|
|
20
|
+
Noul float64 `json:"noul"`
|
|
21
|
+
}
|
|
22
|
+
|
|
23
|
+
// ChoiceAnswer is a selected label and its probabilities.
|
|
24
|
+
type ChoiceAnswer struct {
|
|
25
|
+
Choice string `json:"choice"`
|
|
26
|
+
Confidence float64 `json:"confidence"`
|
|
27
|
+
Probabilities map[string]float64 `json:"probabilities"`
|
|
28
|
+
}
|
|
29
|
+
|
|
30
|
+
// ScoreAnswer is an expected score with its rubric and probabilities.
|
|
31
|
+
type ScoreAnswer struct {
|
|
32
|
+
// Score is the expected score, which may fall between integer rubric levels.
|
|
33
|
+
Score float64 `json:"score"`
|
|
34
|
+
Confidence float64 `json:"confidence"`
|
|
35
|
+
// Legend maps a score to the rubric description that was sent for it.
|
|
36
|
+
Legend map[int]any `json:"legend"`
|
|
37
|
+
// Probabilities maps a score to its probability.
|
|
38
|
+
Probabilities map[int]float64 `json:"probabilities"`
|
|
39
|
+
}
|
|
40
|
+
|
|
41
|
+
// RawAnswer is an answer whose type this package does not know; Raw is the JSON as sent.
|
|
42
|
+
type RawAnswer struct {
|
|
43
|
+
Type QuestionType
|
|
44
|
+
Raw json.RawMessage
|
|
45
|
+
}
|
|
46
|
+
|
|
47
|
+
// AnswerType implements [Answer].
|
|
48
|
+
func (NoulAnswer) AnswerType() QuestionType { return TypeNoul }
|
|
49
|
+
|
|
50
|
+
// AnswerType implements [Answer].
|
|
51
|
+
func (ChoiceAnswer) AnswerType() QuestionType { return TypeChoice }
|
|
52
|
+
|
|
53
|
+
// AnswerType implements [Answer].
|
|
54
|
+
func (ScoreAnswer) AnswerType() QuestionType { return TypeScore }
|
|
55
|
+
|
|
56
|
+
// AnswerType implements [Answer].
|
|
57
|
+
func (a RawAnswer) AnswerType() QuestionType { return a.Type }
|
|
58
|
+
func (NoulAnswer) isAnswer() {}
|
|
59
|
+
func (ChoiceAnswer) isAnswer() {}
|
|
60
|
+
func (ScoreAnswer) isAnswer() {}
|
|
61
|
+
func (RawAnswer) isAnswer() {}
|
|
62
|
+
|
|
63
|
+
// MarshalJSON writes the wire form with its "type".
|
|
64
|
+
func (a NoulAnswer) MarshalJSON() ([]byte, error) {
|
|
65
|
+
type plain NoulAnswer
|
|
66
|
+
return json.Marshal(struct {
|
|
67
|
+
Type QuestionType `json:"type"`
|
|
68
|
+
plain
|
|
69
|
+
}{TypeNoul, plain(a)})
|
|
70
|
+
}
|
|
71
|
+
|
|
72
|
+
// MarshalJSON writes the wire form with its "type".
|
|
73
|
+
func (a ChoiceAnswer) MarshalJSON() ([]byte, error) {
|
|
74
|
+
type plain ChoiceAnswer
|
|
75
|
+
return json.Marshal(struct {
|
|
76
|
+
Type QuestionType `json:"type"`
|
|
77
|
+
plain
|
|
78
|
+
}{TypeChoice, plain(a)})
|
|
79
|
+
}
|
|
80
|
+
|
|
81
|
+
// MarshalJSON writes the wire form with its "type".
|
|
82
|
+
func (a ScoreAnswer) MarshalJSON() ([]byte, error) {
|
|
83
|
+
type plain ScoreAnswer
|
|
84
|
+
return json.Marshal(struct {
|
|
85
|
+
Type QuestionType `json:"type"`
|
|
86
|
+
plain
|
|
87
|
+
}{TypeScore, plain(a)})
|
|
88
|
+
}
|
|
89
|
+
|
|
90
|
+
// MarshalJSON writes the JSON as it was received.
|
|
91
|
+
func (a RawAnswer) MarshalJSON() ([]byte, error) { return a.Raw, nil }
|
|
92
|
+
|
|
93
|
+
// Usage is the token usage of a request.
|
|
94
|
+
type Usage struct {
|
|
95
|
+
InputTokens int `json:"input_tokens"`
|
|
96
|
+
OutputTokens int `json:"output_tokens"`
|
|
97
|
+
}
|
|
98
|
+
|
|
99
|
+
// SystemOneResult is the answers, keyed by question name, with model and usage metadata.
|
|
100
|
+
type SystemOneResult struct {
|
|
101
|
+
// Model is the model that answered the request.
|
|
102
|
+
Model string
|
|
103
|
+
Answers map[string]Answer
|
|
104
|
+
Usage Usage
|
|
105
|
+
}
|
|
106
|
+
|
|
107
|
+
func (r *SystemOneResult) answer(name string) (Answer, error) {
|
|
108
|
+
a, ok := r.Answers[name]
|
|
109
|
+
if !ok {
|
|
110
|
+
return nil, errorf("No answer named %q in the result.", name)
|
|
111
|
+
}
|
|
112
|
+
return a, nil
|
|
113
|
+
}
|
|
114
|
+
|
|
115
|
+
func wrongType(name string, got Answer, want QuestionType) error {
|
|
116
|
+
return errorf("Answer %q is a %s answer, not a %s answer.", name, got.AnswerType(), want)
|
|
117
|
+
}
|
|
118
|
+
|
|
119
|
+
// Noul returns the yes/no answer named name, or an error when it is absent or of
|
|
120
|
+
// another type.
|
|
121
|
+
func (r *SystemOneResult) Noul(name string) (NoulAnswer, error) {
|
|
122
|
+
a, err := r.answer(name)
|
|
123
|
+
if err != nil {
|
|
124
|
+
return NoulAnswer{}, err
|
|
125
|
+
}
|
|
126
|
+
v, ok := a.(NoulAnswer)
|
|
127
|
+
if !ok {
|
|
128
|
+
return NoulAnswer{}, wrongType(name, a, TypeNoul)
|
|
129
|
+
}
|
|
130
|
+
return v, nil
|
|
131
|
+
}
|
|
132
|
+
|
|
133
|
+
// Choice returns the choice answer named name, or an error when it is absent or of
|
|
134
|
+
// another type.
|
|
135
|
+
func (r *SystemOneResult) Choice(name string) (ChoiceAnswer, error) {
|
|
136
|
+
a, err := r.answer(name)
|
|
137
|
+
if err != nil {
|
|
138
|
+
return ChoiceAnswer{}, err
|
|
139
|
+
}
|
|
140
|
+
v, ok := a.(ChoiceAnswer)
|
|
141
|
+
if !ok {
|
|
142
|
+
return ChoiceAnswer{}, wrongType(name, a, TypeChoice)
|
|
143
|
+
}
|
|
144
|
+
return v, nil
|
|
145
|
+
}
|
|
146
|
+
|
|
147
|
+
// Score returns the score answer named name, or an error when it is absent or of
|
|
148
|
+
// another type.
|
|
149
|
+
func (r *SystemOneResult) Score(name string) (ScoreAnswer, error) {
|
|
150
|
+
a, err := r.answer(name)
|
|
151
|
+
if err != nil {
|
|
152
|
+
return ScoreAnswer{}, err
|
|
153
|
+
}
|
|
154
|
+
v, ok := a.(ScoreAnswer)
|
|
155
|
+
if !ok {
|
|
156
|
+
return ScoreAnswer{}, wrongType(name, a, TypeScore)
|
|
157
|
+
}
|
|
158
|
+
return v, nil
|
|
159
|
+
}
|
|
160
|
+
|
|
161
|
+
// UnmarshalJSON decodes the wire form ({model, answers, usage}); each answer is
|
|
162
|
+
// decoded by its "type". Unknown fields are ignored and missing answers are not an error.
|
|
163
|
+
func (r *SystemOneResult) UnmarshalJSON(data []byte) error {
|
|
164
|
+
var wire struct {
|
|
165
|
+
Model string `json:"model"`
|
|
166
|
+
Answers map[string]json.RawMessage `json:"answers"`
|
|
167
|
+
Usage Usage `json:"usage"`
|
|
168
|
+
}
|
|
169
|
+
if err := json.Unmarshal(data, &wire); err != nil {
|
|
170
|
+
return &TypeSafeError{Message: "Unexpected response shape; expected { model, answers, usage }: " + err.Error(), Cause: err}
|
|
171
|
+
}
|
|
172
|
+
out := SystemOneResult{Model: wire.Model, Usage: wire.Usage, Answers: make(map[string]Answer, len(wire.Answers))}
|
|
173
|
+
for name, raw := range wire.Answers {
|
|
174
|
+
a, err := decodeAnswer(raw)
|
|
175
|
+
if err != nil {
|
|
176
|
+
return &TypeSafeError{Message: fmt.Sprintf("Unexpected shape of the answer %q: %v", name, err), Cause: err}
|
|
177
|
+
}
|
|
178
|
+
out.Answers[name] = a
|
|
179
|
+
}
|
|
180
|
+
*r = out
|
|
181
|
+
return nil
|
|
182
|
+
}
|
|
183
|
+
|
|
184
|
+
func decodeAnswer(raw json.RawMessage) (Answer, error) {
|
|
185
|
+
var head struct {
|
|
186
|
+
Type QuestionType `json:"type"`
|
|
187
|
+
}
|
|
188
|
+
if err := json.Unmarshal(raw, &head); err != nil {
|
|
189
|
+
return nil, err
|
|
190
|
+
}
|
|
191
|
+
switch head.Type {
|
|
192
|
+
case TypeNoul:
|
|
193
|
+
var a NoulAnswer
|
|
194
|
+
return a, json.Unmarshal(raw, &a)
|
|
195
|
+
case TypeChoice:
|
|
196
|
+
var a ChoiceAnswer
|
|
197
|
+
return a, json.Unmarshal(raw, &a)
|
|
198
|
+
case TypeScore:
|
|
199
|
+
var a ScoreAnswer
|
|
200
|
+
return a, json.Unmarshal(raw, &a)
|
|
201
|
+
}
|
|
202
|
+
return RawAnswer{Type: head.Type, Raw: append(json.RawMessage(nil), raw...)}, nil
|
|
203
|
+
}
|
|
204
|
+
|
|
205
|
+
// MarshalJSON encodes the wire form.
|
|
206
|
+
func (r SystemOneResult) MarshalJSON() ([]byte, error) {
|
|
207
|
+
names := make([]string, 0, len(r.Answers))
|
|
208
|
+
for n := range r.Answers {
|
|
209
|
+
names = append(names, n)
|
|
210
|
+
}
|
|
211
|
+
sort.Strings(names)
|
|
212
|
+
answers := newObject()
|
|
213
|
+
for _, n := range names {
|
|
214
|
+
answers.value(n, r.Answers[n])
|
|
215
|
+
}
|
|
216
|
+
raw, err := answers.done()
|
|
217
|
+
if err != nil {
|
|
218
|
+
return nil, err
|
|
219
|
+
}
|
|
220
|
+
o := newObject()
|
|
221
|
+
o.value("model", r.Model)
|
|
222
|
+
o.raw("answers", raw)
|
|
223
|
+
o.value("usage", r.Usage)
|
|
224
|
+
return o.done()
|
|
225
|
+
}
|
|
226
|
+
|
|
227
|
+
// ModelCard is the metadata of an available model.
|
|
228
|
+
type ModelCard struct {
|
|
229
|
+
Name string `json:"name"`
|
|
230
|
+
Description string `json:"description"`
|
|
231
|
+
ReleaseDate string `json:"release_date"`
|
|
232
|
+
}
|
|
233
|
+
|
|
234
|
+
// SystemOneRequest is the state and named questions for SystemOne.
|
|
235
|
+
type SystemOneRequest struct {
|
|
236
|
+
// State is text, a JSON object or array, or null.
|
|
237
|
+
State Entry
|
|
238
|
+
// Questions must be non-empty.
|
|
239
|
+
Questions Questions
|
|
240
|
+
// Model overrides the client's default model when non-empty.
|
|
241
|
+
Model string
|
|
242
|
+
// Extra fields are forwarded at the top level of the request body, including nulls.
|
|
243
|
+
// The names state, questions and model are reserved.
|
|
244
|
+
Extra map[string]any
|
|
245
|
+
}
|
|
246
|
+
|
|
247
|
+
// MarshalJSON writes the request body: state, questions, model, then the extra fields.
|
|
248
|
+
func (r SystemOneRequest) MarshalJSON() ([]byte, error) {
|
|
249
|
+
o := newObject()
|
|
250
|
+
o.entry("state", r.State)
|
|
251
|
+
o.value("questions", r.Questions)
|
|
252
|
+
if r.Model != "" {
|
|
253
|
+
o.value("model", r.Model)
|
|
254
|
+
}
|
|
255
|
+
names := make([]string, 0, len(r.Extra))
|
|
256
|
+
for n := range r.Extra {
|
|
257
|
+
names = append(names, n)
|
|
258
|
+
}
|
|
259
|
+
sort.Strings(names)
|
|
260
|
+
for _, n := range names {
|
|
261
|
+
switch n {
|
|
262
|
+
case "state", "questions", "model":
|
|
263
|
+
return nil, errorf("The extra request field %q collides with a request field of the same name.", n)
|
|
264
|
+
}
|
|
265
|
+
o.value(n, r.Extra[n])
|
|
266
|
+
}
|
|
267
|
+
return o.done()
|
|
268
|
+
}
|