Files
scriptorium/internal/usecase/runner.go

372 lines
10 KiB
Go

package usecase
import (
"context"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"strings"
"time"
"gitea.maximumdirect.net/eric/scriptorium/internal/artifact"
"gitea.maximumdirect.net/eric/scriptorium/internal/domain"
"gitea.maximumdirect.net/eric/scriptorium/internal/llm"
"gitea.maximumdirect.net/eric/scriptorium/internal/profile"
"gitea.maximumdirect.net/eric/scriptorium/internal/prompt"
"gitea.maximumdirect.net/eric/scriptorium/internal/promptdef"
"gitea.maximumdirect.net/eric/scriptorium/internal/validate"
)
var (
ErrInvalidRequest = errors.New("invalid run request")
ErrProfileLoad = errors.New("failed to load prompt definition")
ErrArtifactLoad = errors.New("failed to load artifact")
ErrPromptRender = errors.New("failed to render prompt")
ErrLLMGenerate = errors.New("failed to generate output")
ErrValidation = errors.New("failed to validate output")
)
// Runner executes the Scriptorium core use case.
type Runner struct {
promptDefs promptdef.Repository
profiles profile.Repository
artifacts artifact.Reader
renderer prompt.Renderer
llm llm.Client
validator validate.Validator
repairer OutputRepairer
}
func NewRunner(
promptDefs promptdef.Repository,
profiles profile.Repository,
artifacts artifact.Reader,
renderer prompt.Renderer,
llmClient llm.Client,
validator validate.Validator,
) *Runner {
return NewRunnerWithRepairer(promptDefs, profiles, artifacts, renderer, llmClient, validator, nil)
}
func NewRunnerWithRepairer(
promptDefs promptdef.Repository,
profiles profile.Repository,
artifacts artifact.Reader,
renderer prompt.Renderer,
llmClient llm.Client,
validator validate.Validator,
repairer OutputRepairer,
) *Runner {
return &Runner{
promptDefs: promptDefs,
profiles: profiles,
artifacts: artifacts,
renderer: renderer,
llm: llmClient,
validator: validator,
repairer: repairer,
}
}
func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunResult, error) {
if strings.TrimSpace(req.PromptID) == "" {
return nil, fmt.Errorf("%w: prompt id is required", ErrInvalidRequest)
}
runID, err := newRunID()
if err != nil {
return nil, fmt.Errorf("failed to create run id: %w", err)
}
start := time.Now().UTC()
def, err := r.promptDefs.GetPromptDefinition(ctx, req.PromptID, req.PromptVersion)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrProfileLoad, err)
}
promptDefinitionHash, err := hashPromptDefinition(def)
if err != nil {
return nil, fmt.Errorf("%w: failed to hash prompt definition: %v", ErrProfileLoad, err)
}
selectedProfileID := strings.TrimSpace(req.ProfileID)
if selectedProfileID == "" {
selectedProfileID = strings.TrimSpace(def.DefaultProfile)
}
if selectedProfileID == "" {
return nil, fmt.Errorf("%w: profile id is required either in request or prompt default_profile", ErrInvalidRequest)
}
execProfile, err := r.profiles.GetProfile(ctx, selectedProfileID)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrProfileLoad, err)
}
effectiveModel := mergeExecutionTarget(executionProfileToTarget(execProfile), req.Execution)
if strings.TrimSpace(effectiveModel.Endpoint) == "" {
return nil, fmt.Errorf("%w: execution endpoint is required", ErrInvalidRequest)
}
if strings.TrimSpace(effectiveModel.Model) == "" {
return nil, fmt.Errorf("%w: execution model is required", ErrInvalidRequest)
}
effectiveContract := resolveOutputContract(def, req.Validation)
resolvedInputs := make(map[string]*domain.Artifact, len(req.Inputs))
inputHashes := make(map[string]string, len(req.Inputs))
for name, ref := range req.Inputs {
art, readErr := r.artifacts.Read(ctx, ref)
if readErr != nil {
return nil, fmt.Errorf("%w: input %q: %w", ErrArtifactLoad, name, readErr)
}
if art.Name == "" {
art.Name = name
}
resolvedInputs[name] = art
inputHashes[name] = art.Hash
}
renderedPrompt, err := r.renderer.Render(ctx, def, resolvedInputs, req.Vars)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrPromptRender, err)
}
renderedPromptHash := hashRenderedPrompt(*renderedPrompt)
genResp, err := r.llm.Generate(ctx, domain.GenerateRequest{
Prompt: *renderedPrompt,
Target: effectiveModel,
})
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrLLMGenerate, err)
}
outputArtifact := buildOutputArtifact(genResp.Content, effectiveContract.Format)
validationResult, err := r.validateOutput(ctx, &outputArtifact, effectiveContract, 0)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrValidation, err)
}
if r.shouldAttemptRepair(effectiveContract, validationResult) {
attemptsUsed := 0
for attemptsUsed < effectiveContract.RepairAttempts && validationResult.Status == domain.ValidationFailed {
attemptsUsed++
repairResp, repairErr := r.repairer.Repair(ctx, RepairRequest{
PreviousOutput: genResp.Content,
ValidationErrors: validationResult.Errors,
Target: effectiveModel,
Attempt: attemptsUsed,
MaxAttempts: effectiveContract.RepairAttempts,
Mode: effectiveContract.ValidationMode,
})
if repairErr != nil {
return nil, fmt.Errorf("%w: %w", ErrValidation, repairErr)
}
if repairResp == nil {
return nil, fmt.Errorf("%w: repairer returned nil response", ErrValidation)
}
genResp = repairResp
outputArtifact = buildOutputArtifact(genResp.Content, effectiveContract.Format)
validationResult, err = r.validateOutput(ctx, &outputArtifact, effectiveContract, attemptsUsed)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrValidation, err)
}
}
}
end := time.Now().UTC()
return &domain.RunResult{
RunID: runID,
Artifact: outputArtifact,
RawOutput: genResp.Content,
Validation: validationResult,
PromptID: def.ID,
PromptVersion: def.Version,
PromptHash: promptDefinitionHash,
RenderedPromptHash: renderedPromptHash,
SelectedProfileID: selectedProfileID,
ModelName: effectiveModel.Model,
Endpoint: effectiveModel.Endpoint,
EffectiveModelParams: effectiveModel,
InputHashes: inputHashes,
Usage: genResp.Usage,
StartTime: start,
EndTime: end,
Duration: end.Sub(start),
}, nil
}
func (r *Runner) validateOutput(ctx context.Context, artifact *domain.Artifact, contract domain.OutputContract, attemptsUsed int) (domain.ValidationResult, error) {
if r.validator == nil || contract.ValidationMode == domain.ValidationNone {
return domain.ValidationResult{
Status: domain.ValidationSkipped,
Mode: contract.ValidationMode,
SchemaPath: contract.SchemaPath,
RepairAttempts: attemptsUsed,
IsValid: true,
}, nil
}
res, err := r.validator.Validate(ctx, artifact, contract)
if err != nil {
return domain.ValidationResult{}, err
}
res.RepairAttempts = attemptsUsed
return res, nil
}
func (r *Runner) shouldAttemptRepair(contract domain.OutputContract, validationResult domain.ValidationResult) bool {
if r.repairer == nil {
return false
}
if contract.RepairAttempts <= 0 {
return false
}
if validationResult.Status != domain.ValidationFailed {
return false
}
return contract.ValidationMode == domain.ValidationJSON || contract.ValidationMode == domain.ValidationJSONSchema
}
func mergeExecutionTarget(base domain.ExecutionTarget, override *domain.ExecutionTarget) domain.ExecutionTarget {
if override == nil {
return base
}
out := base
if override.Endpoint != "" {
out.Endpoint = override.Endpoint
}
if override.Model != "" {
out.Model = override.Model
}
if override.Temperature != 0 {
out.Temperature = override.Temperature
}
if override.MaxTokens != 0 {
out.MaxTokens = override.MaxTokens
}
if override.TopP != 0 {
out.TopP = override.TopP
}
if override.TimeoutSeconds != 0 {
out.TimeoutSeconds = override.TimeoutSeconds
}
if strings.TrimSpace(override.ReasoningEffort) != "" {
out.ReasoningEffort = override.ReasoningEffort
}
if strings.TrimSpace(override.APIKeyEnv) != "" {
out.APIKeyEnv = override.APIKeyEnv
}
if len(override.ExtraParams) > 0 {
cp := make(map[string]string, len(override.ExtraParams))
for k, v := range override.ExtraParams {
cp[k] = v
}
out.ExtraParams = cp
}
return out
}
func executionProfileToTarget(p *domain.ExecutionProfile) domain.ExecutionTarget {
if p == nil {
return domain.ExecutionTarget{}
}
cp := map[string]string(nil)
if len(p.ExtraParams) > 0 {
cp = make(map[string]string, len(p.ExtraParams))
for k, v := range p.ExtraParams {
cp[k] = v
}
}
return domain.ExecutionTarget{
Endpoint: p.Endpoint,
Model: p.Model,
Temperature: p.Temperature,
MaxTokens: p.MaxTokens,
TopP: p.TopP,
TimeoutSeconds: p.TimeoutSeconds,
ReasoningEffort: p.ReasoningEffort,
APIKeyEnv: p.APIKeyEnv,
ExtraParams: cp,
}
}
func resolveOutputContract(def *domain.PromptDefinition, override *domain.OutputContract) domain.OutputContract {
contract := def.Validation
if contract.Format == "" {
contract.Format = def.OutputFormat
}
if override != nil {
contract = *override
}
if contract.Format == "" {
contract.Format = domain.FormatText
}
return contract
}
func hashRenderedPrompt(p domain.RenderedPrompt) string {
var b strings.Builder
for _, msg := range p.Messages {
b.WriteString(msg.Role)
b.WriteByte('\n')
b.WriteString(msg.Content)
b.WriteString("\n---\n")
}
h := sha256.Sum256([]byte(b.String()))
return hex.EncodeToString(h[:])
}
func buildOutputArtifact(content string, format domain.OutputFormat) domain.Artifact {
body := []byte(content)
hash := sha256.Sum256(body)
contentType := "text/plain"
switch format {
case domain.FormatMarkdown:
contentType = "text/markdown"
case domain.FormatJSON:
contentType = "application/json"
}
return domain.Artifact{
Name: "output",
ContentType: contentType,
Body: body,
Size: int64(len(body)),
Hash: hex.EncodeToString(hash[:]),
}
}
func hashPromptDefinition(def *domain.PromptDefinition) (string, error) {
b, err := json.Marshal(def)
if err != nil {
return "", err
}
sum := sha256.Sum256(b)
return hex.EncodeToString(sum[:]), nil
}
func newRunID() (string, error) {
var b [16]byte
if _, err := rand.Read(b[:]); err != nil {
return "", err
}
// UUID v4 (RFC 4122 variant).
b[6] = (b[6] & 0x0f) | 0x40
b[8] = (b[8] & 0x3f) | 0x80
return fmt.Sprintf("%08x-%04x-%04x-%04x-%012x",
b[0:4],
b[4:6],
b[6:8],
b[8:10],
b[10:16],
), nil
}