150 lines
4.6 KiB
Go
150 lines
4.6 KiB
Go
package profile
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"strings"
|
|
|
|
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
|
"gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue"
|
|
)
|
|
|
|
const maximumProfileChainLength = 32
|
|
|
|
type resolvingRepository struct {
|
|
source Repository
|
|
}
|
|
|
|
// NewResolvingRepository resolves inherited profile definitions from source.
|
|
func NewResolvingRepository(source Repository) Repository {
|
|
return &resolvingRepository{source: source}
|
|
}
|
|
|
|
func (r *resolvingRepository) GetProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) {
|
|
if r == nil || r.source == nil {
|
|
return nil, fmt.Errorf("%w: profile repository is required", ErrInvalidProfile)
|
|
}
|
|
requestedID := strings.TrimSpace(id)
|
|
if requestedID == "" {
|
|
return nil, fmt.Errorf("%w: profile id is required", ErrInvalidProfile)
|
|
}
|
|
if err := ctx.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
profile, err := r.source.GetProfile(ctx, requestedID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if profile == nil {
|
|
return nil, fmt.Errorf("%w: selected profile %q is nil", ErrInvalidProfile, requestedID)
|
|
}
|
|
|
|
chain := []*domain.ExecutionProfile{profile}
|
|
chainIDs := []string{requestedID}
|
|
visited := map[string]struct{}{requestedID: {}}
|
|
current := profile
|
|
|
|
for {
|
|
baseID := strings.TrimSpace(current.BaseProfileID)
|
|
if baseID == "" {
|
|
break
|
|
}
|
|
if err := ctx.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
if _, seen := visited[baseID]; seen {
|
|
return nil, fmt.Errorf("%w: profile inheritance cycle %s", ErrInvalidProfile, joinProfileChain(chainIDs, baseID))
|
|
}
|
|
if len(chain) >= maximumProfileChainLength {
|
|
return nil, fmt.Errorf("%w: profile inheritance chain exceeds %d profiles: %s", ErrInvalidProfile, maximumProfileChainLength, joinProfileChain(chainIDs, baseID))
|
|
}
|
|
|
|
base, err := r.source.GetProfile(ctx, baseID)
|
|
if err != nil {
|
|
if errors.Is(err, ErrProfileNotFound) {
|
|
return nil, fmt.Errorf("%w: base profile %q is missing in chain %s", ErrInvalidProfile, baseID, joinProfileChain(chainIDs, baseID))
|
|
}
|
|
return nil, fmt.Errorf("%w: failed to load base profile %q in chain %s: %w", ErrInvalidProfile, baseID, joinProfileChain(chainIDs, baseID), err)
|
|
}
|
|
if base == nil {
|
|
return nil, fmt.Errorf("%w: base profile %q is nil in chain %s", ErrInvalidProfile, baseID, joinProfileChain(chainIDs, baseID))
|
|
}
|
|
|
|
chain = append(chain, base)
|
|
chainIDs = append(chainIDs, baseID)
|
|
visited[baseID] = struct{}{}
|
|
current = base
|
|
}
|
|
|
|
resolved, err := mergeProfileChain(chain)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("%w: resolved profile chain %s: %w", ErrInvalidProfile, strings.Join(chainIDs, " -> "), err)
|
|
}
|
|
if err := validateResolvedProfile(resolved); err != nil {
|
|
return nil, fmt.Errorf("%w: resolved profile chain %s: %w", ErrInvalidProfile, strings.Join(chainIDs, " -> "), err)
|
|
}
|
|
return resolved, nil
|
|
}
|
|
|
|
func joinProfileChain(chain []string, next string) string {
|
|
return strings.Join(append(append([]string(nil), chain...), next), " -> ")
|
|
}
|
|
|
|
func mergeProfileChain(chain []*domain.ExecutionProfile) (*domain.ExecutionProfile, error) {
|
|
resolved := &domain.ExecutionProfile{ID: chain[0].ID}
|
|
for index := len(chain) - 1; index >= 0; index-- {
|
|
definition := chain[index]
|
|
if strings.TrimSpace(definition.BackendID) != "" {
|
|
resolved.BackendID = definition.BackendID
|
|
}
|
|
if strings.TrimSpace(definition.Endpoint) != "" {
|
|
resolved.Endpoint = definition.Endpoint
|
|
}
|
|
if strings.TrimSpace(definition.Model) != "" {
|
|
resolved.Model = definition.Model
|
|
}
|
|
if definition.Temperature != 0 {
|
|
resolved.Temperature = definition.Temperature
|
|
}
|
|
if definition.MaxTokens != 0 {
|
|
resolved.MaxTokens = definition.MaxTokens
|
|
}
|
|
if definition.TopP != 0 {
|
|
resolved.TopP = definition.TopP
|
|
}
|
|
if definition.TimeoutSeconds != 0 {
|
|
resolved.TimeoutSeconds = definition.TimeoutSeconds
|
|
}
|
|
if strings.TrimSpace(definition.ServiceTier) != "" {
|
|
resolved.ServiceTier = definition.ServiceTier
|
|
}
|
|
if strings.TrimSpace(definition.ReasoningEffort) != "" {
|
|
resolved.ReasoningEffort = definition.ReasoningEffort
|
|
}
|
|
if strings.TrimSpace(definition.APIKeyEnv) != "" {
|
|
resolved.APIKeyEnv = definition.APIKeyEnv
|
|
}
|
|
resolved.APIKeyRequired = resolved.APIKeyRequired || definition.APIKeyRequired
|
|
if len(definition.ExtraParams) != 0 {
|
|
extraParams, err := jsonvalue.CopyMap(definition.ExtraParams)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
resolved.ExtraParams = extraParams
|
|
}
|
|
}
|
|
resolved.ID = chain[0].ID
|
|
resolved.BaseProfileID = ""
|
|
return resolved, nil
|
|
}
|
|
|
|
func validateResolvedProfile(profile *domain.ExecutionProfile) error {
|
|
if profile == nil {
|
|
return errors.New("resolved profile is required")
|
|
}
|
|
profile.BaseProfileID = ""
|
|
return NormalizeAndValidateDefinition(profile)
|
|
}
|