Add public artifact reader support
This commit is contained in:
155
engine_test.go
155
engine_test.go
@@ -835,11 +835,132 @@ func TestMissingCredentialsFailClearlyWhenProfileRequiresAuth(t *testing.T) {
|
||||
if !errors.Is(err, scriptorium.ErrInvalidRequest) {
|
||||
t.Fatalf("expected invalid request for missing credentials, got %v", err)
|
||||
}
|
||||
if !errors.Is(err, scriptorium.ErrAPIKeyEnvMissing) {
|
||||
t.Fatalf("expected missing credential environment error, got %v", err)
|
||||
}
|
||||
if err == nil || !strings.Contains(err.Error(), missingEnv) {
|
||||
t.Fatalf("expected missing env name in error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithArtifactReaderRejectsNilReader(t *testing.T) {
|
||||
_, err := scriptorium.NewEngine(contractConfig(frameworkSchemaDir), scriptorium.WithArtifactReader(nil))
|
||||
if !errors.Is(err, scriptorium.ErrInvalidConfig) {
|
||||
t.Fatalf("expected ErrInvalidConfig, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArtifactReaderReceivesPublicReferenceAndPreparesArtifact(t *testing.T) {
|
||||
reader := &recordingArtifactReader{
|
||||
artifact: &scriptorium.Artifact{
|
||||
ContentType: "text/plain",
|
||||
Body: []byte("Reader-supplied transcript."),
|
||||
URI: "reader://transcript",
|
||||
Size: int64(len("Reader-supplied transcript.")),
|
||||
Hash: "reader-transcript-hash",
|
||||
},
|
||||
}
|
||||
engine := newArtifactReaderEngine(t, reader)
|
||||
|
||||
ref := scriptorium.ArtifactRef{
|
||||
Type: scriptorium.ArtifactRefInline,
|
||||
URI: "reader://transcript",
|
||||
Body: "request body",
|
||||
}
|
||||
prepared, err := engine.Prepare(context.Background(), scriptorium.RunRequest{
|
||||
PromptID: "artifact-reader",
|
||||
Inputs: map[string]scriptorium.ArtifactRef{
|
||||
"transcript": ref,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare with artifact reader: %v", err)
|
||||
}
|
||||
if len(reader.refs) != 1 || !reflect.DeepEqual(reader.refs[0], ref) {
|
||||
t.Fatalf("reader received %#v, want %#v", reader.refs, ref)
|
||||
}
|
||||
if prepared.InputHashes["transcript"] != "reader-transcript-hash" {
|
||||
t.Fatalf("unexpected input hash: %#v", prepared.InputHashes)
|
||||
}
|
||||
if len(prepared.Messages) != 1 || !strings.Contains(prepared.Messages[0].Content, "Reader-supplied transcript.") {
|
||||
t.Fatalf("prepared prompt omitted reader artifact: %#v", prepared.Messages)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArtifactReaderFailuresPreserveArtifactLoadErrors(t *testing.T) {
|
||||
readerErr := errors.New("artifact reader failed")
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
ctx context.Context
|
||||
reader *recordingArtifactReader
|
||||
wantNested error
|
||||
}{
|
||||
{
|
||||
name: "reader error",
|
||||
ctx: context.Background(),
|
||||
reader: &recordingArtifactReader{err: readerErr},
|
||||
wantNested: readerErr,
|
||||
},
|
||||
{
|
||||
name: "nil artifact",
|
||||
ctx: context.Background(),
|
||||
reader: &recordingArtifactReader{},
|
||||
},
|
||||
{
|
||||
name: "reader cancellation",
|
||||
ctx: context.Background(),
|
||||
reader: &recordingArtifactReader{read: func(context.Context, scriptorium.ArtifactRef) (*scriptorium.Artifact, error) {
|
||||
return nil, context.Canceled
|
||||
}},
|
||||
wantNested: context.Canceled,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
engine := newArtifactReaderEngine(t, tc.reader)
|
||||
_, err := engine.Prepare(tc.ctx, scriptorium.RunRequest{
|
||||
PromptID: "artifact-reader",
|
||||
Inputs: map[string]scriptorium.ArtifactRef{
|
||||
"transcript": scriptorium.Inline("input"),
|
||||
},
|
||||
})
|
||||
if !errors.Is(err, scriptorium.ErrArtifactLoad) {
|
||||
t.Fatalf("expected ErrArtifactLoad, got %v", err)
|
||||
}
|
||||
if tc.wantNested != nil && !errors.Is(err, tc.wantNested) {
|
||||
t.Fatalf("expected nested %v, got %v", tc.wantNested, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrepareWithoutProfileMatchesSpecificPublicError(t *testing.T) {
|
||||
promptDir := t.TempDir()
|
||||
writePublicPromptFile(t, promptDir, "profile-required", "")
|
||||
engine, err := scriptorium.NewEngine(scriptorium.Config{
|
||||
PromptDir: promptDir,
|
||||
SchemaDir: frameworkSchemaDir,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine: %v", err)
|
||||
}
|
||||
|
||||
_, err = engine.Prepare(context.Background(), scriptorium.RunRequest{
|
||||
PromptID: "profile-required",
|
||||
Inputs: map[string]scriptorium.ArtifactRef{
|
||||
"transcript": scriptorium.Inline("input"),
|
||||
},
|
||||
})
|
||||
if !errors.Is(err, scriptorium.ErrInvalidRequest) {
|
||||
t.Fatalf("expected ErrInvalidRequest, got %v", err)
|
||||
}
|
||||
if !errors.Is(err, scriptorium.ErrProfileRequired) {
|
||||
t.Fatalf("expected ErrProfileRequired, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunValidationFailureReturnsResult(t *testing.T) {
|
||||
fake := &fakeLLMClient{
|
||||
response: &scriptorium.GenerateResponse{Content: ""},
|
||||
@@ -2207,6 +2328,22 @@ func newContractEngineWithOptions(t *testing.T, schemaDir string, opts ...script
|
||||
return engine
|
||||
}
|
||||
|
||||
func newArtifactReaderEngine(t *testing.T, reader scriptorium.ArtifactReader) *scriptorium.Engine {
|
||||
t.Helper()
|
||||
|
||||
promptDir := t.TempDir()
|
||||
writePublicPromptFile(t, promptDir, "artifact-reader", frameworkFastProfileID)
|
||||
engine, err := scriptorium.NewEngine(scriptorium.Config{
|
||||
PromptDir: promptDir,
|
||||
ProfileDir: frameworkProfileDir,
|
||||
SchemaDir: frameworkSchemaDir,
|
||||
}, scriptorium.WithArtifactReader(reader))
|
||||
if err != nil {
|
||||
t.Fatalf("construct engine with artifact reader: %v", err)
|
||||
}
|
||||
return engine
|
||||
}
|
||||
|
||||
func contractConfig(schemaDir string) scriptorium.Config {
|
||||
return scriptorium.Config{
|
||||
PromptDir: frameworkPromptDir,
|
||||
@@ -2363,6 +2500,24 @@ type fakeLLMClient struct {
|
||||
requests []scriptorium.GenerateRequest
|
||||
}
|
||||
|
||||
type recordingArtifactReader struct {
|
||||
artifact *scriptorium.Artifact
|
||||
err error
|
||||
refs []scriptorium.ArtifactRef
|
||||
read func(context.Context, scriptorium.ArtifactRef) (*scriptorium.Artifact, error)
|
||||
}
|
||||
|
||||
func (r *recordingArtifactReader) Read(ctx context.Context, ref scriptorium.ArtifactRef) (*scriptorium.Artifact, error) {
|
||||
r.refs = append(r.refs, ref)
|
||||
if r.read != nil {
|
||||
return r.read(ctx, ref)
|
||||
}
|
||||
if r.err != nil {
|
||||
return nil, r.err
|
||||
}
|
||||
return r.artifact, nil
|
||||
}
|
||||
|
||||
type roundTripFunc func(*http.Request) (*http.Response, error)
|
||||
|
||||
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
|
||||
Reference in New Issue
Block a user