Files
promptkit/internal/usecase/runner.go

748 lines
22 KiB
Go

package usecase
import (
"context"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"os"
"strings"
"time"
"unicode"
"gitea.maximumdirect.net/eric/promptkit/internal/artifact"
"gitea.maximumdirect.net/eric/promptkit/internal/capacity"
"gitea.maximumdirect.net/eric/promptkit/internal/defaults"
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
"gitea.maximumdirect.net/eric/promptkit/internal/llm"
"gitea.maximumdirect.net/eric/promptkit/internal/profile"
"gitea.maximumdirect.net/eric/promptkit/internal/prompt"
"gitea.maximumdirect.net/eric/promptkit/internal/promptdef"
"gitea.maximumdirect.net/eric/promptkit/internal/validate"
)
var (
ErrInvalidRequest = errors.New("invalid run request")
ErrProfileRequired = errors.New("profile selection is required")
ErrAPIKeyEnvMissing = errors.New("api_key_env points to an unset environment variable")
ErrAPIKeyRequired = errors.New("api key is required")
ErrPromptLoad = errors.New("failed to load prompt definition")
ErrProfileLoad = errors.New("failed to load execution profile")
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 Promptkit core use case.
type Runner struct {
promptDefs promptdef.Repository
profiles profile.Repository
backends BackendResolver
artifacts artifact.Reader
renderer prompt.Renderer
llm llm.Client
validator validate.Validator
repairer OutputRepairer
admitter RunAdmitter
}
// BackendResolver resolves one normalized backend ID.
type BackendResolver interface {
GetBackend(string) (domain.Backend, error)
}
// RunAdmitter reserves capacity for one resolved backend run.
type RunAdmitter interface {
Admit(context.Context, string) (func(), error)
}
type preparationState struct {
definition *domain.PromptDefinition
directSessionID string
promptDefinitionHash string
selectedProfileID string
effectiveModel domain.ExecutionTarget
targetPresence domain.ExecutionTargetPresence
effectiveContract domain.OutputContract
start time.Time
}
func NewRunner(
promptDefs promptdef.Repository,
profiles profile.Repository,
backends BackendResolver,
artifacts artifact.Reader,
renderer prompt.Renderer,
llmClient llm.Client,
validator validate.Validator,
admitter RunAdmitter,
) *Runner {
return NewRunnerWithRepairer(
promptDefs,
profiles,
backends,
artifacts,
renderer,
llmClient,
validator,
nil,
admitter,
)
}
func NewRunnerWithRepairer(
promptDefs promptdef.Repository,
profiles profile.Repository,
backends BackendResolver,
artifacts artifact.Reader,
renderer prompt.Renderer,
llmClient llm.Client,
validator validate.Validator,
repairer OutputRepairer,
admitter RunAdmitter,
) *Runner {
return &Runner{
promptDefs: promptDefs,
profiles: profiles,
backends: backends,
artifacts: artifacts,
renderer: renderer,
llm: llmClient,
validator: validator,
repairer: repairer,
admitter: admitter,
}
}
func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunResult, error) {
runID, err := newRunID()
if err != nil {
return nil, fmt.Errorf("failed to create run id: %w", err)
}
start := time.Now().UTC()
state, err := r.resolvePreparation(ctx, req, time.Now().UTC())
if err != nil {
return nil, err
}
release, err := r.admitRun(ctx, state.effectiveModel.BackendID)
if err != nil {
return nil, err
}
defer release()
prepared, err := r.completePreparation(ctx, req, state)
if err != nil {
return nil, err
}
directAPIKey := state.effectiveModel.APIKey
return r.executePreparedRun(ctx, prepared, directAPIKey, runID, start, func(
ctx context.Context,
artifact *domain.Artifact,
attemptsUsed int,
) (domain.ValidationResult, error) {
return r.validateOutput(ctx, artifact, prepared.OutputContract, attemptsUsed)
})
}
type preparedValidationFunc func(
context.Context,
*domain.Artifact,
int,
) (domain.ValidationResult, error)
func (r *Runner) executePreparedRun(
ctx context.Context,
prepared *domain.PreparedRun,
directAPIKey string,
runID string,
start time.Time,
validateArtifact preparedValidationFunc,
) (*domain.RunResult, error) {
executionTarget := prepared.EffectiveModelParams
executionTarget.APIKey = directAPIKey
genResp, err := r.llm.Generate(ctx, domain.GenerateRequest{
Prompt: domain.RenderedPrompt{SessionID: prepared.SessionID, Messages: prepared.Messages},
Target: executionTarget,
TargetPresence: prepared.TargetPresence,
StructuredOutput: prepared.StructuredOutput,
})
if err != nil {
if errors.Is(err, llm.ErrInvalidRequest) {
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
}
return nil, fmt.Errorf("%w: %w", ErrLLMGenerate, err)
}
outputArtifact := buildOutputArtifact(genResp.Content, prepared.OutputContract.Format)
validationResult, err := validateArtifact(ctx, &outputArtifact, 0)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrValidation, err)
}
if r.shouldAttemptRepair(prepared.OutputContract, validationResult) {
attemptsUsed := 0
for attemptsUsed < prepared.OutputContract.RepairAttempts && validationResult.Status == domain.ValidationFailed {
attemptsUsed++
repairResp, repairErr := r.repairer.Repair(ctx, RepairRequest{
PreviousOutput: genResp.Content,
ValidationErrors: validationResult.Errors,
SessionID: prepared.SessionID,
Target: executionTarget,
StructuredOutput: prepared.StructuredOutput,
Attempt: attemptsUsed,
MaxAttempts: prepared.OutputContract.RepairAttempts,
Mode: prepared.OutputContract.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, prepared.OutputContract.Format)
validationResult, err = validateArtifact(ctx, &outputArtifact, attemptsUsed)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrValidation, err)
}
}
}
end := time.Now().UTC()
executionTarget.APIKey = ""
return &domain.RunResult{
RunID: runID,
Artifact: outputArtifact,
RawOutput: genResp.Content,
Validation: validationResult,
PromptID: prepared.PromptID,
PromptVersion: prepared.PromptVersion,
PromptHash: prepared.PromptHash,
SessionID: prepared.SessionID,
RenderedPromptHash: prepared.RenderedPromptHash,
SelectedProfileID: prepared.SelectedProfileID,
SelectedBackendID: prepared.SelectedBackendID,
ModelName: prepared.EffectiveModelParams.Model,
Endpoint: prepared.EffectiveModelParams.Endpoint,
EffectiveModelParams: executionTarget,
InputHashes: prepared.InputHashes,
Usage: genResp.Usage,
StartTime: start,
EndTime: end,
Duration: end.Sub(start),
}, nil
}
func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.PreparedRun, error) {
state, err := r.resolvePreparation(ctx, req, time.Now().UTC())
if err != nil {
return nil, err
}
return r.completePreparation(ctx, req, state)
}
func (r *Runner) resolvePreparation(
ctx context.Context,
req domain.RunRequest,
start time.Time,
) (*preparationState, error) {
if strings.TrimSpace(req.PromptID) == "" {
return nil, fmt.Errorf("%w: prompt id is required", ErrInvalidRequest)
}
directSessionID, err := domain.NormalizeSessionID(req.SessionID)
if err != nil {
return nil, fmt.Errorf("%w: session_id: %v", ErrInvalidRequest, err)
}
promptSelection, err := r.resolvePromptDefinition(ctx, req.PromptID, req.PromptVersion)
if err != nil {
return nil, err
}
def := promptSelection.definition
promptDefinitionHash := promptSelection.hash
selectedProfileID := strings.TrimSpace(req.ProfileID)
if selectedProfileID == "" {
selectedProfileID = strings.TrimSpace(def.DefaultProfile)
}
if selectedProfileID == "" {
return nil, fmt.Errorf("%w: %w: profile id is required either in request or prompt default_profile", ErrInvalidRequest, ErrProfileRequired)
}
selection, err := r.resolveProfileSelection(ctx, selectedProfileID)
if err != nil {
return nil, err
}
effectiveModel, targetPresence, err := resolveExecutionTarget(selection.backend, selection.profile, req.Execution)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
}
effectiveModel.APIKey = req.APIKey
if err := validateResolvedExecutionTarget(effectiveModel); err != nil {
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
}
if err := validateAPIKey(effectiveModel.APIKeyEnv, effectiveModel.APIKey, effectiveModel.APIKeyRequired); err != nil {
return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err)
}
effectiveContract := resolveOutputContract(def, req.Validation)
return &preparationState{
definition: def,
directSessionID: directSessionID,
promptDefinitionHash: promptDefinitionHash,
selectedProfileID: selection.id,
effectiveModel: effectiveModel,
targetPresence: targetPresence,
effectiveContract: effectiveContract,
start: start,
}, nil
}
func (r *Runner) completePreparation(
ctx context.Context,
req domain.RunRequest,
state *preparationState,
) (*domain.PreparedRun, error) {
structuredOutput, err := r.resolveStructuredOutput(
ctx,
state.definition,
state.effectiveContract,
)
if err != nil {
return nil, err
}
return r.completePreparationWithStructuredOutput(ctx, req, state, structuredOutput)
}
func (r *Runner) completePreparationWithStructuredOutput(
ctx context.Context,
req domain.RunRequest,
state *preparationState,
structuredOutput *domain.StructuredOutputSpec,
) (*domain.PreparedRun, error) {
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
}
definitionToRender := state.definition
if state.directSessionID != "" {
definitionCopy := *state.definition
definitionCopy.SessionID = ""
definitionToRender = &definitionCopy
}
renderedPrompt, err := r.renderer.Render(ctx, definitionToRender, resolvedInputs, req.Vars)
if err != nil {
return nil, fmt.Errorf("%w: %w", ErrPromptRender, err)
}
if state.directSessionID != "" {
renderedPrompt.SessionID = state.directSessionID
}
end := time.Now().UTC()
effectiveModel := state.effectiveModel
effectiveModel.APIKey = ""
return &domain.PreparedRun{
PromptID: state.definition.ID,
PromptVersion: state.definition.Version,
PromptHash: state.promptDefinitionHash,
SelectedProfileID: state.selectedProfileID,
SelectedBackendID: state.effectiveModel.BackendID,
EffectiveModelParams: effectiveModel,
TargetPresence: state.targetPresence,
OutputContract: state.effectiveContract,
StructuredOutput: structuredOutput,
InputHashes: inputHashes,
SessionID: renderedPrompt.SessionID,
RenderedPromptHash: hashRenderedPrompt(*renderedPrompt),
Messages: renderedPrompt.Messages,
StartTime: state.start,
EndTime: end,
DurationMS: end.Sub(state.start).Milliseconds(),
}, nil
}
func (r *Runner) admitRun(ctx context.Context, backendID string) (func(), error) {
if r.admitter == nil {
return func() {}, nil
}
release, err := r.admitter.Admit(ctx, backendID)
if err != nil {
if errors.Is(err, capacity.ErrCapacityExceeded) && strings.TrimSpace(backendID) != "" {
return nil, &CapacityError{BackendID: backendID}
}
return nil, err
}
return release, nil
}
func (r *Runner) resolveStructuredOutput(ctx context.Context, def *domain.PromptDefinition, contract domain.OutputContract) (*domain.StructuredOutputSpec, error) {
if contract.ValidationMode != domain.ValidationJSONSchema {
return nil, nil
}
loader, ok := r.validator.(validate.SchemaDocumentLoader)
if !ok || loader == nil {
return nil, fmt.Errorf("%w: json_schema output requires schema document loader", ErrValidation)
}
schemaDoc, err := loader.LoadSchemaDocument(ctx, contract.SchemaPath)
if err != nil {
return nil, fmt.Errorf("%w: failed to load json schema for structured output: %v", ErrValidation, err)
}
return structuredOutputSpec(def, schemaDoc), nil
}
func structuredOutputSpec(def *domain.PromptDefinition, schemaDocument any) *domain.StructuredOutputSpec {
return &domain.StructuredOutputSpec{
Type: domain.StructuredOutputJSONSchema,
JSONSchema: &domain.StructuredOutputJSONSpec{
Name: deriveStructuredSchemaName(def.ID, def.Version),
Strict: true,
Schema: schemaDocument,
},
}
}
func deriveStructuredSchemaName(promptID string, promptVersion string) string {
raw := strings.TrimSpace(promptID)
if v := strings.TrimSpace(promptVersion); v != "" {
if raw == "" {
raw = v
} else {
raw = raw + "_" + v
}
}
var b strings.Builder
for _, r := range raw {
if unicode.IsLetter(r) || unicode.IsDigit(r) || r == '_' || r == '-' {
b.WriteRune(r)
} else {
b.WriteRune('_')
}
}
name := strings.Trim(b.String(), "_-")
if name == "" {
return "promptkit_schema"
}
return name
}
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 {
out := base
if strings.TrimSpace(override.BackendID) != "" {
out.BackendID = override.BackendID
}
if strings.TrimSpace(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.ServiceTier) != "" {
out.ServiceTier = override.ServiceTier
}
if strings.TrimSpace(override.ReasoningEffort) != "" {
out.ReasoningEffort = override.ReasoningEffort
}
if strings.TrimSpace(override.APIKeyEnv) != "" {
out.APIKeyEnv = override.APIKeyEnv
}
if override.APIKeyRequired {
out.APIKeyRequired = true
out.APIKeyEnv = ""
}
if len(override.ExtraParams) > 0 {
out.ExtraParams = copyExtraParams(override.ExtraParams)
}
return out
}
func mergeExecutionTargetOverride(base domain.ExecutionTarget, override domain.ExecutionTargetOverride) (domain.ExecutionTarget, domain.ExecutionTargetPresence, error) {
out := base
var presence domain.ExecutionTargetPresence
if override.Endpoint != "" {
out.Endpoint = override.Endpoint
}
if override.Model != "" {
out.Model = override.Model
}
if override.Temperature != nil {
if *override.Temperature < 0 || *override.Temperature > 2 {
return domain.ExecutionTarget{}, domain.ExecutionTargetPresence{}, errors.New("temperature must be between 0 and 2")
}
out.Temperature = *override.Temperature
presence.Temperature = true
}
if override.MaxTokens != nil {
if *override.MaxTokens < 0 {
return domain.ExecutionTarget{}, domain.ExecutionTargetPresence{}, errors.New("max_tokens must be greater than or equal to 0")
}
out.MaxTokens = *override.MaxTokens
presence.MaxTokens = true
}
if override.TopP != nil {
if *override.TopP < 0 || *override.TopP > 1 {
return domain.ExecutionTarget{}, domain.ExecutionTargetPresence{}, errors.New("top_p must be between 0 and 1")
}
out.TopP = *override.TopP
presence.TopP = true
}
if override.TimeoutSeconds != nil {
if *override.TimeoutSeconds < 0 {
return domain.ExecutionTarget{}, domain.ExecutionTargetPresence{}, errors.New("timeout_seconds must be greater than or equal to 0")
}
out.TimeoutSeconds = *override.TimeoutSeconds
presence.TimeoutSeconds = true
}
if strings.TrimSpace(override.ServiceTier) != "" {
out.ServiceTier = override.ServiceTier
}
if override.ReasoningEffort != nil {
out.ReasoningEffort = strings.TrimSpace(*override.ReasoningEffort)
}
if strings.TrimSpace(override.APIKeyEnv) != "" {
out.APIKeyEnv = override.APIKeyEnv
}
if len(override.ExtraParams) > 0 {
out.ExtraParams = copyExtraParams(override.ExtraParams)
}
return out, presence, nil
}
func resolveExecutionTarget(backendValue *domain.Backend, profileValue *domain.ExecutionProfile, override *domain.ExecutionTargetOverride) (domain.ExecutionTarget, domain.ExecutionTargetPresence, error) {
out := defaults.ExecutionTargetDefault()
out = mergeExecutionTarget(out, backendToTarget(backendValue))
out = mergeExecutionTarget(out, executionProfileToTarget(profileValue))
var presence domain.ExecutionTargetPresence
if override != nil {
var err error
out, presence, err = mergeExecutionTargetOverride(out, *override)
if err != nil {
return domain.ExecutionTarget{}, domain.ExecutionTargetPresence{}, err
}
}
return out, presence, nil
}
func validateAPIKey(apiKeyEnv string, apiKey string, apiKeyRequired bool) error {
if strings.TrimSpace(apiKey) != "" {
return nil
}
envName := strings.TrimSpace(apiKeyEnv)
if envName == "" {
if apiKeyRequired {
return ErrAPIKeyRequired
}
return nil
}
if strings.TrimSpace(os.Getenv(envName)) == "" {
return fmt.Errorf("%w: api key environment variable %q is not set", ErrAPIKeyEnvMissing, envName)
}
return nil
}
func executionProfileToTarget(p *domain.ExecutionProfile) domain.ExecutionTarget {
if p == nil {
return domain.ExecutionTarget{}
}
endpoint := p.Endpoint
if strings.TrimSpace(endpoint) == "" {
endpoint = ""
}
return domain.ExecutionTarget{
BackendID: p.BackendID,
Endpoint: endpoint,
Model: p.Model,
Temperature: p.Temperature,
MaxTokens: p.MaxTokens,
TopP: p.TopP,
TimeoutSeconds: p.TimeoutSeconds,
ServiceTier: p.ServiceTier,
ReasoningEffort: p.ReasoningEffort,
APIKeyEnv: p.APIKeyEnv,
APIKeyRequired: p.APIKeyRequired,
ExtraParams: copyExtraParams(p.ExtraParams),
}
}
func backendToTarget(value *domain.Backend) domain.ExecutionTarget {
if value == nil {
return domain.ExecutionTarget{}
}
return domain.ExecutionTarget{
BackendID: value.ID,
Endpoint: value.Endpoint,
APIKeyEnv: value.APIKeyEnv,
ExtraParams: copyExtraParams(value.ExtraParams),
}
}
func copyExtraParams(src map[string]any) map[string]any {
if len(src) == 0 {
return nil
}
cp := make(map[string]any, len(src))
for k, v := range src {
cp[k] = v
}
return 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
if p.SessionID != "" {
b.WriteString("session_id=")
b.WriteString(p.SessionID)
b.WriteString("\n---\n")
}
for _, msg := range p.Messages {
b.WriteString(msg.Role)
b.WriteByte('\n')
b.WriteString(msg.Content)
if msg.CacheControl != nil {
b.WriteString("\ncache_control.type=")
b.WriteString(string(msg.CacheControl.Type))
if msg.CacheControl.TTL != "" {
b.WriteString("\ncache_control.ttl=")
b.WriteString(msg.CacheControl.TTL)
}
}
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 := defaults.ContentTypeTextPlain
switch format {
case domain.FormatMarkdown:
contentType = defaults.ContentTypeTextMarkdown
case domain.FormatJSON:
contentType = defaults.ContentTypeApplicationJSON
}
return domain.Artifact{
Name: defaults.OutputArtifactName,
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
}