Refactor: split prompt definition from execution settings and migrate run contracts to prompt_* + execution_target

This commit is contained in:
2026-05-05 10:09:31 -05:00
parent fdfd8641f5
commit a633c67538
28 changed files with 712 additions and 1021 deletions

View File

@@ -7,6 +7,7 @@ import (
"errors"
"fmt"
"io"
"os"
"net/http"
"net/url"
"strings"
@@ -25,7 +26,6 @@ var (
type OpenAICompatibleConfig struct {
BaseURL string
APIKey string
Model string
Timeout time.Duration
HTTPClient *http.Client
@@ -33,7 +33,6 @@ type OpenAICompatibleConfig struct {
type OpenAICompatibleClient struct {
baseURL string
apiKey string
defaultModel string
timeout time.Duration
httpClient *http.Client
@@ -64,7 +63,6 @@ func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleCli
return &OpenAICompatibleClient{
baseURL: strings.TrimRight(baseURL, "/"),
apiKey: cfg.APIKey,
defaultModel: cfg.Model,
timeout: timeout,
httpClient: client,
@@ -125,8 +123,12 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera
return nil, fmt.Errorf("%w: failed to create request: %v", ErrRequestFailed, err)
}
httpReq.Header.Set("Content-Type", "application/json")
if strings.TrimSpace(c.apiKey) != "" {
httpReq.Header.Set("Authorization", "Bearer "+c.apiKey)
if envName := strings.TrimSpace(req.Target.APIKeyEnv); envName != "" {
apiKey := strings.TrimSpace(os.Getenv(envName))
if apiKey == "" {
return nil, fmt.Errorf("%w: api key environment variable %q is not set", ErrInvalidRequest, envName)
}
httpReq.Header.Set("Authorization", "Bearer "+apiKey)
}
effectiveTimeout := c.timeout

View File

@@ -44,23 +44,24 @@ func TestOpenAICompatibleClientGenerateSuccess(t *testing.T) {
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: ts.URL + "/v1",
APIKey: "secret-key",
Timeout: 2 * time.Second,
})
if err != nil {
t.Fatalf("unexpected constructor error: %v", err)
}
t.Setenv("SCRIPTORIUM_TEST_API_KEY", "secret-key")
resp, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{
{Role: "system", Content: "You are helpful."},
{Role: "user", Content: "Say hello"},
}},
Target: domain.ModelTarget{
Target: domain.ExecutionTarget{
Model: "gpt-test",
Temperature: 0.4,
MaxTokens: 123,
TopP: 0.7,
APIKeyEnv: "SCRIPTORIUM_TEST_API_KEY",
},
})
if err != nil {
@@ -110,7 +111,7 @@ func TestOpenAICompatibleClientNoAuthorizationHeaderWhenNoAPIKey(t *testing.T) {
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ModelTarget{Model: "model"},
Target: domain.ExecutionTarget{Model: "model"},
})
if err != nil {
t.Fatalf("expected no error, got %v", err)
@@ -120,6 +121,29 @@ func TestOpenAICompatibleClientNoAuthorizationHeaderWhenNoAPIKey(t *testing.T) {
}
}
func TestOpenAICompatibleClientAPIKeyEnvMissing(t *testing.T) {
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`))
}))
defer ts.Close()
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1", Model: "model"})
if err != nil {
t.Fatal(err)
}
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ExecutionTarget{APIKeyEnv: "SCRIPTORIUM_MISSING_KEY"},
})
if err == nil {
t.Fatal("expected missing API key env error")
}
if !errors.Is(err, ErrInvalidRequest) {
t.Fatalf("expected ErrInvalidRequest, got %v", err)
}
}
func TestOpenAICompatibleClientModelFallbackFromConfig(t *testing.T) {
gotModel := ""
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
@@ -139,7 +163,7 @@ func TestOpenAICompatibleClientModelFallbackFromConfig(t *testing.T) {
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ModelTarget{},
Target: domain.ExecutionTarget{},
})
if err != nil {
t.Fatalf("expected no error, got %v", err)
@@ -175,7 +199,7 @@ func TestOpenAICompatibleClientEndpointOverride(t *testing.T) {
resp, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ModelTarget{Endpoint: overrideServer.URL + "/v1"},
Target: domain.ExecutionTarget{Endpoint: overrideServer.URL + "/v1"},
})
if err != nil {
t.Fatalf("expected no error, got %v", err)
@@ -306,7 +330,7 @@ func TestOpenAICompatibleClientRequestTimeoutOverride(t *testing.T) {
resp, err := client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ModelTarget{TimeoutSeconds: 1},
Target: domain.ExecutionTarget{TimeoutSeconds: 1},
})
if err != nil {
t.Fatalf("expected request-level timeout override to succeed, got %v", err)
@@ -327,7 +351,7 @@ func TestOpenAICompatibleClientNegativeTimeoutRejected(t *testing.T) {
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ModelTarget{TimeoutSeconds: -1},
Target: domain.ExecutionTarget{TimeoutSeconds: -1},
})
if err == nil {
t.Fatal("expected invalid request error")
@@ -348,7 +372,7 @@ func TestOpenAICompatibleClientAllowsEmptyConfiguredBaseURL(t *testing.T) {
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ModelTarget{Endpoint: "http://localhost:9999/v1"},
Target: domain.ExecutionTarget{Endpoint: "http://localhost:9999/v1"},
})
if err == nil {
t.Fatal("expected request failure due to unreachable endpoint")
@@ -369,7 +393,7 @@ func TestOpenAICompatibleClientRequiresEndpointWhenUnsetEverywhere(t *testing.T)
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}},
Target: domain.ModelTarget{},
Target: domain.ExecutionTarget{},
})
if err == nil {
t.Fatal("expected endpoint-required error")