Files
notarius/internal/framework/llm/scriptorium_client.go

268 lines
7.7 KiB
Go

package llm
import (
"context"
"encoding/json"
"fmt"
"net/http"
"regexp"
"sort"
"strings"
"sync"
"time"
"gitea.maximumdirect.net/eric/notarius/internal/core/artifacts"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
"gitea.maximumdirect.net/eric/scriptorium"
)
const scriptoriumProviderName = "scriptorium"
type ScriptoriumClientConfig struct {
ProfileDir string
ProfileFile string
Assets *AssetRegistry
Timeout time.Duration
HTTPClient *http.Client
EngineOptions []scriptorium.Option
Recorder *LLMProfileRecorder
}
type ScriptoriumClient struct {
engine *scriptorium.Engine
recorder *LLMProfileRecorder
}
type LLMProfileRecorder struct {
mu sync.Mutex
profiles map[string]artifacts.LLMProfileManifest
}
var _ contracts.StructuredLLMClient = (*ScriptoriumClient)(nil)
var _ contracts.LLMProfileManifestProvider = (*ScriptoriumClient)(nil)
func NewScriptoriumClient(cfg ScriptoriumClientConfig) (*ScriptoriumClient, error) {
if cfg.Assets == nil {
return nil, fmt.Errorf("scriptorium client assets must not be nil")
}
if strings.TrimSpace(cfg.ProfileDir) != "" && strings.TrimSpace(cfg.ProfileFile) != "" {
return nil, fmt.Errorf("scriptorium profile_dir and profile_file are mutually exclusive")
}
options, err := cfg.Assets.ScriptoriumOptions()
if err != nil {
return nil, err
}
if profileFile := strings.TrimSpace(cfg.ProfileFile); profileFile != "" {
options = append(options, scriptorium.WithProfileFile(profileFile))
}
options = append(options, cfg.EngineOptions...)
engine, err := scriptorium.NewEngine(scriptorium.Config{
ProfileDir: strings.TrimSpace(cfg.ProfileDir),
Timeout: cfg.Timeout,
HTTPClient: cfg.HTTPClient,
}, options...)
if err != nil {
return nil, fmt.Errorf("create Scriptorium engine: %w", err)
}
recorder := cfg.Recorder
if recorder == nil {
recorder = NewLLMProfileRecorder()
}
return &ScriptoriumClient{
engine: engine,
recorder: recorder,
}, nil
}
func (c *ScriptoriumClient) CompleteStructured(ctx context.Context, req contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) {
if c == nil {
return contracts.StructuredCompletionResponse{}, fmt.Errorf("scriptorium client must not be nil")
}
if c.engine == nil {
return contracts.StructuredCompletionResponse{}, fmt.Errorf("scriptorium client engine must not be nil")
}
if err := validateOutputTarget(out); err != nil {
return contracts.StructuredCompletionResponse{}, err
}
promptID := strings.TrimSpace(req.PromptID)
if promptID == "" {
return contracts.StructuredCompletionResponse{}, fmt.Errorf("structured completion prompt_id must not be empty")
}
runReq := scriptorium.RunRequest{
PromptID: promptID,
PromptVersion: strings.TrimSpace(req.PromptVersion),
ProfileID: strings.TrimSpace(req.ProfileID),
Inputs: scriptoriumInputs(req.Inputs),
Vars: scriptoriumVars(req),
Metadata: scriptoriumMetadata(req),
}
result, err := c.engine.Run(ctx, runReq)
if err != nil {
if ctxErr := ctx.Err(); ctxErr != nil {
return contracts.StructuredCompletionResponse{}, ctxErr
}
return contracts.StructuredCompletionResponse{}, fmt.Errorf("run Scriptorium prompt %q: %w", promptID, redactScriptoriumError(err))
}
if result == nil {
return contracts.StructuredCompletionResponse{}, fmt.Errorf("run Scriptorium prompt %q: empty result", promptID)
}
if result.Validation.Status == scriptorium.ValidationFailed || !result.Validation.IsValid {
return contracts.StructuredCompletionResponse{}, fmt.Errorf("run Scriptorium prompt %q: validation failed: %s", promptID, strings.Join(result.Validation.Errors, "; "))
}
content := result.Artifact.Body
if len(content) == 0 {
content = []byte(result.RawOutput)
}
if len(strings.TrimSpace(string(content))) == 0 {
return contracts.StructuredCompletionResponse{}, fmt.Errorf("run Scriptorium prompt %q: empty structured output", promptID)
}
if err := json.Unmarshal(content, out); err != nil {
return contracts.StructuredCompletionResponse{}, fmt.Errorf("decode Scriptorium structured output for prompt %q: %w", promptID, err)
}
profile := artifacts.LLMProfileManifest{
ID: strings.TrimSpace(result.SelectedProfileID),
Provider: scriptoriumProviderName,
Model: firstNonEmpty(result.ModelName, result.EffectiveModelParams.Model),
}
if c.recorder != nil {
c.recorder.Record(profile)
}
return contracts.StructuredCompletionResponse{
Content: append(json.RawMessage(nil), content...),
Provider: profile.Provider,
Model: profile.Model,
ProfileID: profile.ID,
PromptTokens: result.Usage.PromptTokens,
CompletionTokens: result.Usage.CompletionTokens,
TotalTokens: result.Usage.TotalTokens,
}, nil
}
func (c *ScriptoriumClient) LLMProfileManifests() []artifacts.LLMProfileManifest {
if c == nil || c.recorder == nil {
return nil
}
return c.recorder.Manifests()
}
func NewLLMProfileRecorder() *LLMProfileRecorder {
return &LLMProfileRecorder{profiles: map[string]artifacts.LLMProfileManifest{}}
}
func (r *LLMProfileRecorder) Record(profile artifacts.LLMProfileManifest) {
if r == nil {
return
}
profile.ID = strings.TrimSpace(profile.ID)
profile.Provider = strings.TrimSpace(profile.Provider)
profile.Model = strings.TrimSpace(profile.Model)
key := profile.ID + "\x00" + profile.Provider + "\x00" + profile.Model
r.mu.Lock()
defer r.mu.Unlock()
if r.profiles == nil {
r.profiles = map[string]artifacts.LLMProfileManifest{}
}
r.profiles[key] = profile
}
func (r *LLMProfileRecorder) Manifests() []artifacts.LLMProfileManifest {
if r == nil {
return nil
}
r.mu.Lock()
defer r.mu.Unlock()
if len(r.profiles) == 0 {
return nil
}
keys := make([]string, 0, len(r.profiles))
for key := range r.profiles {
keys = append(keys, key)
}
sort.Strings(keys)
out := make([]artifacts.LLMProfileManifest, 0, len(keys))
for _, key := range keys {
out = append(out, r.profiles[key])
}
return out
}
func scriptoriumInputs(inputs contracts.LLMInputSet) map[string]scriptorium.ArtifactRef {
if len(inputs) == 0 {
return nil
}
out := make(map[string]scriptorium.ArtifactRef, len(inputs))
for key, material := range inputs {
name := strings.TrimSpace(key)
if name == "" {
name = strings.TrimSpace(material.Name)
}
if name == "" {
continue
}
body := string(material.Content)
if body == "" {
body = " "
}
if origin := strings.TrimSpace(material.OriginURI); origin != "" {
out[name] = scriptorium.InlineWithURI(origin, body)
} else {
out[name] = scriptorium.Inline(body)
}
}
return out
}
func scriptoriumVars(req contracts.StructuredCompletionRequest) map[string]string {
vars := make(map[string]string, len(req.Vars)+1)
for key, value := range req.Vars {
name := strings.TrimSpace(key)
if name == "" || value == nil {
continue
}
vars[name] = fmt.Sprint(value)
}
if sessionID := strings.TrimSpace(req.SessionID); sessionID != "" {
vars["session_id"] = sessionID
}
if len(vars) == 0 {
return nil
}
return vars
}
func scriptoriumMetadata(req contracts.StructuredCompletionRequest) map[string]string {
metadata := map[string]string{}
if stageName := strings.TrimSpace(req.StageName); stageName != "" {
metadata["stage_name"] = stageName
}
if len(metadata) == 0 {
return nil
}
return metadata
}
var bearerTokenPattern = regexp.MustCompile(`(?i)Bearer\s+[A-Za-z0-9._~+/=-]+`)
func redactScriptoriumError(err error) error {
if err == nil {
return nil
}
return redactedProviderError{err: err}
}
type redactedProviderError struct {
err error
}
func (e redactedProviderError) Error() string {
return bearerTokenPattern.ReplaceAllString(e.err.Error(), "Bearer "+secretReplacement)
}
func (e redactedProviderError) Unwrap() error {
return e.err
}