268 lines
6.4 KiB
Go
268 lines
6.4 KiB
Go
package profile
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"io/fs"
|
|
"os"
|
|
"path"
|
|
"strings"
|
|
|
|
"gitea.maximumdirect.net/eric/scriptorium/internal/domain"
|
|
"gopkg.in/yaml.v3"
|
|
)
|
|
|
|
var (
|
|
ErrProfileNotFound = errors.New("execution profile not found")
|
|
ErrInvalidYAML = errors.New("invalid YAML format")
|
|
ErrInvalidProfile = errors.New("invalid execution profile configuration")
|
|
ErrRawAPIKeyNotAllowed = errors.New("raw api_key is not allowed; use api_key_env")
|
|
)
|
|
|
|
type filesystemRepository struct {
|
|
dir string
|
|
}
|
|
|
|
func NewFilesystemRepository(dir string) Repository {
|
|
return &filesystemRepository{dir: dir}
|
|
}
|
|
|
|
func (r *filesystemRepository) GetProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) {
|
|
return loadProfile(ctx, os.DirFS(r.dir), ".", id)
|
|
}
|
|
|
|
type fsRepository struct {
|
|
fsys fs.FS
|
|
root string
|
|
}
|
|
|
|
func NewFSRepository(fsys fs.FS, root string) Repository {
|
|
return &fsRepository{fsys: fsys, root: root}
|
|
}
|
|
|
|
func (r *fsRepository) GetProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) {
|
|
return loadProfile(ctx, r.fsys, r.root, id)
|
|
}
|
|
|
|
type overlayRepository struct {
|
|
primary Repository
|
|
fallback Repository
|
|
}
|
|
|
|
func NewOverlayRepository(primary, fallback Repository) Repository {
|
|
return &overlayRepository{primary: primary, fallback: fallback}
|
|
}
|
|
|
|
func (r *overlayRepository) GetProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) {
|
|
if r.primary != nil {
|
|
prof, err := r.primary.GetProfile(ctx, id)
|
|
if err == nil {
|
|
return prof, nil
|
|
}
|
|
if !errors.Is(err, ErrProfileNotFound) {
|
|
return nil, err
|
|
}
|
|
}
|
|
if r.fallback == nil {
|
|
return nil, ErrProfileNotFound
|
|
}
|
|
return r.fallback.GetProfile(ctx, id)
|
|
}
|
|
|
|
func loadProfile(ctx context.Context, fsys fs.FS, root string, id string) (*domain.ExecutionProfile, error) {
|
|
if strings.TrimSpace(id) == "" {
|
|
return nil, fmt.Errorf("%w: profile id is required", ErrInvalidProfile)
|
|
}
|
|
if fsys == nil {
|
|
return nil, fmt.Errorf("failed to read profile directory: filesystem is nil")
|
|
}
|
|
|
|
files, err := findProfileYAMLFiles(ctx, fsys, root)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to read profile directory: %w", err)
|
|
}
|
|
|
|
var matches []profileMatch
|
|
for _, fullPath := range files {
|
|
select {
|
|
case <-ctx.Done():
|
|
return nil, ctx.Err()
|
|
default:
|
|
}
|
|
|
|
relPath := displayPath(root, fullPath)
|
|
fileMatch := profileFileStem(path.Base(fullPath)) == id
|
|
data, err := fs.ReadFile(fsys, fullPath)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to read profile file %s: %w", relPath, err)
|
|
}
|
|
metadata := readProfileFileMetadata(data)
|
|
idMatch := fileMatch || metadata.id == id
|
|
if metadata.hasRawAPIKey {
|
|
if idMatch {
|
|
return nil, fmt.Errorf("%w: %s", ErrRawAPIKeyNotAllowed, relPath)
|
|
}
|
|
continue
|
|
}
|
|
|
|
var prof domain.ExecutionProfile
|
|
decoder := yaml.NewDecoder(bytes.NewReader(data))
|
|
decoder.KnownFields(true)
|
|
if err := decoder.Decode(&prof); err != nil {
|
|
if idMatch {
|
|
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidYAML, relPath, err)
|
|
}
|
|
continue
|
|
}
|
|
|
|
if prof.ID != id {
|
|
continue
|
|
}
|
|
if err := validateProfile(&prof); err != nil {
|
|
if errors.Is(err, ErrRawAPIKeyNotAllowed) {
|
|
return nil, fmt.Errorf("%w: %s", err, relPath)
|
|
}
|
|
return nil, fmt.Errorf("%w: %s: %v", ErrInvalidProfile, relPath, err)
|
|
}
|
|
matches = append(matches, profileMatch{
|
|
profile: &prof,
|
|
path: relPath,
|
|
})
|
|
}
|
|
|
|
if len(matches) > 1 {
|
|
paths := make([]string, 0, len(matches))
|
|
for _, match := range matches {
|
|
paths = append(paths, match.path)
|
|
}
|
|
return nil, fmt.Errorf("%w: duplicate execution profile id %q found in: %s", ErrInvalidProfile, id, strings.Join(paths, ", "))
|
|
}
|
|
|
|
if len(matches) == 1 {
|
|
return matches[0].profile, nil
|
|
}
|
|
|
|
return nil, ErrProfileNotFound
|
|
}
|
|
|
|
func findProfileYAMLFiles(ctx context.Context, fsys fs.FS, root string) ([]string, error) {
|
|
cleanRoot := cleanFSRoot(root)
|
|
var files []string
|
|
err := fs.WalkDir(fsys, cleanRoot, func(name string, d fs.DirEntry, err error) error {
|
|
if err != nil {
|
|
return err
|
|
}
|
|
select {
|
|
case <-ctx.Done():
|
|
return ctx.Err()
|
|
default:
|
|
}
|
|
if d.IsDir() {
|
|
return nil
|
|
}
|
|
if !isProfileYAMLFile(d.Name()) {
|
|
return nil
|
|
}
|
|
files = append(files, name)
|
|
return nil
|
|
})
|
|
return files, err
|
|
}
|
|
|
|
func cleanFSRoot(root string) string {
|
|
root = strings.TrimSpace(root)
|
|
if root == "" || root == "." {
|
|
return "."
|
|
}
|
|
return path.Clean(root)
|
|
}
|
|
|
|
func displayPath(root string, name string) string {
|
|
cleanRoot := cleanFSRoot(root)
|
|
cleanName := path.Clean(name)
|
|
if cleanRoot == "." {
|
|
return cleanName
|
|
}
|
|
prefix := strings.TrimSuffix(cleanRoot, "/") + "/"
|
|
if strings.HasPrefix(cleanName, prefix) {
|
|
return strings.TrimPrefix(cleanName, prefix)
|
|
}
|
|
return cleanName
|
|
}
|
|
|
|
func profileFileStem(name string) string {
|
|
name = strings.TrimSuffix(name, ".yaml")
|
|
name = strings.TrimSuffix(name, ".yml")
|
|
return name
|
|
}
|
|
|
|
func isProfileYAMLFile(name string) bool {
|
|
return strings.HasSuffix(name, ".yaml") || strings.HasSuffix(name, ".yml")
|
|
}
|
|
|
|
type profileMatch struct {
|
|
profile *domain.ExecutionProfile
|
|
path string
|
|
}
|
|
|
|
type profileFileMetadata struct {
|
|
id string
|
|
hasRawAPIKey bool
|
|
}
|
|
|
|
func readProfileFileMetadata(data []byte) profileFileMetadata {
|
|
var node yaml.Node
|
|
if err := yaml.NewDecoder(bytes.NewReader(data)).Decode(&node); err != nil {
|
|
return profileFileMetadata{}
|
|
}
|
|
if node.Kind != yaml.DocumentNode || len(node.Content) == 0 {
|
|
return profileFileMetadata{}
|
|
}
|
|
mapping := node.Content[0]
|
|
if mapping.Kind != yaml.MappingNode {
|
|
return profileFileMetadata{}
|
|
}
|
|
|
|
var metadata profileFileMetadata
|
|
for i := 0; i+1 < len(mapping.Content); i += 2 {
|
|
key := mapping.Content[i]
|
|
value := mapping.Content[i+1]
|
|
switch key.Value {
|
|
case "id":
|
|
metadata.id = strings.TrimSpace(value.Value)
|
|
case "api_key":
|
|
metadata.hasRawAPIKey = true
|
|
}
|
|
}
|
|
return metadata
|
|
}
|
|
|
|
func validateProfile(p *domain.ExecutionProfile) error {
|
|
if strings.TrimSpace(p.ID) == "" {
|
|
return errors.New("id is required")
|
|
}
|
|
if strings.TrimSpace(p.Endpoint) == "" {
|
|
return errors.New("endpoint is required")
|
|
}
|
|
if strings.TrimSpace(p.Model) == "" {
|
|
return errors.New("model is required")
|
|
}
|
|
|
|
if p.Temperature < 0 || p.Temperature > 2 {
|
|
return errors.New("temperature must be between 0 and 2")
|
|
}
|
|
if p.MaxTokens < 0 {
|
|
return errors.New("max_tokens must be greater than or equal to 0")
|
|
}
|
|
if p.TopP < 0 || p.TopP > 1 {
|
|
return errors.New("top_p must be between 0 and 1")
|
|
}
|
|
if p.TimeoutSeconds < 0 {
|
|
return errors.New("timeout_seconds must be greater than or equal to 0")
|
|
}
|
|
|
|
return nil
|
|
}
|