Files

472 lines
17 KiB
Go

// Package modelmigration creates a reviewable v2 installation candidate from legacy model files.
package modelmigration
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strconv"
"strings"
"github.com/aritmolab/thothii/tools/tht/internal/config"
"github.com/compose-spec/compose-go/v2/dotenv"
"gopkg.in/yaml.v3"
)
const maxMigrationInputBytes = 1 << 20
// Request contains the facts that do not exist unambiguously in the three legacy model sources.
type Request struct {
InstallationPath string
OutputPath string
SessionDefault string
EmbeddingID string
EmbeddingDimensions int
}
type legacyDescriptor struct {
SchemaVersion int `yaml:"schemaVersion,omitempty"`
Profile string `yaml:"profile"`
ProjectDirectory string `yaml:"projectDirectory"`
EnvFile string `yaml:"envFile"`
WorkspaceRepository workspaceRepository `yaml:"workspaceRepository"`
Authentication authenticationDescriptor `yaml:"authentication"`
MetadataGeneration legacyMetadataCatalog `yaml:"metadataGeneration"`
Overrides []string `yaml:"overrides,omitempty"`
}
type workspaceRepository struct {
Remote string `yaml:"remote"`
Branch string `yaml:"branch"`
Access string `yaml:"access"`
}
type authenticationDescriptor struct {
ConfigDirectory string `yaml:"configDirectory"`
RuntimeProjection *runtimeProjection `yaml:"runtimeProjection,omitempty"`
}
type runtimeProjection struct {
Directory string `yaml:"directory"`
UID uint32 `yaml:"uid"`
GID uint32 `yaml:"gid"`
}
type legacyMetadataCatalog struct {
Default string `yaml:"default"`
Models []legacyMetadataModel `yaml:"models"`
}
type legacyMetadataModel struct {
ID string `yaml:"id"`
Label string `yaml:"label"`
LiteLLM legacyLiteLLM `yaml:"litellm"`
APIKeyEnv string `yaml:"apiKeyEnv"`
}
type legacyLiteLLM struct {
Provider string `yaml:"provider"`
Model string `yaml:"model"`
DisableThinking bool `yaml:"disableThinking"`
Endpoint *config.ModelEndpoint `yaml:"endpoint"`
}
type piCatalog struct {
Providers map[string]piProvider `json:"providers"`
}
type piProvider struct {
BaseURL string `json:"baseUrl"`
API string `json:"api"`
APIKey string `json:"apiKey"`
Models []piModel `json:"models"`
}
type piModel struct {
ID string `json:"id"`
Name string `json:"name"`
Reasoning bool `json:"reasoning"`
Input []string `json:"input,omitempty"`
Cost *config.ModelCost `json:"cost,omitempty"`
ContextWindow int `json:"contextWindow"`
MaxTokens int `json:"maxTokens"`
Compat *config.ModelCompatibility `json:"compat,omitempty"`
}
type piSettings struct {
DefaultProjectTrust string `json:"defaultProjectTrust"`
EnabledModels []string `json:"enabledModels"`
}
type candidateDescriptor struct {
SchemaVersion int `yaml:"schemaVersion"`
Profile string `yaml:"profile"`
ProjectDirectory string `yaml:"projectDirectory"`
EnvFile string `yaml:"envFile"`
WorkspaceRepository workspaceRepository `yaml:"workspaceRepository"`
ModelCatalog config.ModelCatalog `yaml:"modelCatalog"`
Authentication authenticationDescriptor `yaml:"authentication"`
Overrides []string `yaml:"overrides,omitempty"`
}
// Run reads every legacy source, reports ambiguity without publishing, and otherwise writes one
// new candidate. It never modifies or removes the legacy files.
func Run(request Request) error {
if !filepath.IsAbs(request.InstallationPath) || filepath.Base(request.InstallationPath) != "thothii-installation.yaml" {
return errors.New("installation migration requires an absolute thothii-installation.yaml path")
}
outputExtension := strings.ToLower(filepath.Ext(request.OutputPath))
if !filepath.IsAbs(request.OutputPath) || outputExtension != ".yaml" && outputExtension != ".yml" ||
filepath.Clean(request.OutputPath) == filepath.Clean(request.InstallationPath) {
return errors.New("installation migration output must be a different absolute YAML path")
}
if request.SessionDefault == "" || request.EmbeddingID == "" || request.EmbeddingDimensions <= 0 {
return errors.New("installation migration requires session default, embedding id, and positive embedding dimensions")
}
if _, err := os.Lstat(request.OutputPath); err == nil {
return errors.New("installation migration output already exists")
} else if !os.IsNotExist(err) {
return errors.New("installation migration output is unavailable")
}
var legacy legacyDescriptor
if err := decodeYAML(request.InstallationPath, &legacy); err != nil {
return fmt.Errorf("legacy installation: %w", err)
}
if legacy.SchemaVersion != 0 {
return errors.New("legacy installation: schemaVersion must be absent")
}
var piModels piCatalog
if err := decodeJSON(filepath.Join(legacy.ProjectDirectory, "deploy", "pi", "models.json"), &piModels); err != nil {
return fmt.Errorf("deploy/pi/models.json: %w", err)
}
var settings piSettings
if err := decodeJSON(filepath.Join(legacy.ProjectDirectory, "deploy", "pi", "settings.json"), &settings); err != nil {
return fmt.Errorf("deploy/pi/settings.json: %w", err)
}
catalog, err := reconcileCatalog(legacy, piModels, settings, request)
if err != nil {
return err
}
environment, err := parseEnvironment(legacy.EnvFile)
if err != nil {
return errors.New("legacy installation envFile is unavailable or invalid")
}
if err := catalog.Validate(environment); err != nil {
return fmt.Errorf("modelCatalog candidate: %w", err)
}
candidate := candidateDescriptor{
SchemaVersion: 2, Profile: legacy.Profile, ProjectDirectory: legacy.ProjectDirectory,
EnvFile: legacy.EnvFile, WorkspaceRepository: legacy.WorkspaceRepository,
ModelCatalog: catalog, Authentication: legacy.Authentication, Overrides: legacy.Overrides,
}
contents, err := yaml.Marshal(candidate)
if err != nil {
return errors.New("installation migration candidate could not be encoded")
}
return publishCandidate(request.OutputPath, contents)
}
func reconcileCatalog(legacy legacyDescriptor, piModels piCatalog, settings piSettings, request Request) (config.ModelCatalog, error) {
catalog := config.ModelCatalog{
Defaults: config.ModelCatalogDefaults{Interaction: request.SessionDefault},
Embedding: config.ModelCatalogEmbedding{ID: request.EmbeddingID, Dimensions: request.EmbeddingDimensions},
Providers: make(map[string]config.ModelProvider),
}
enabled := make(map[string]struct{}, len(settings.EnabledModels))
for index, canonical := range settings.EnabledModels {
providerID, modelID, ok := splitCanonical(canonical)
if !ok {
return catalog, fmt.Errorf("deploy/pi/settings.json enabledModels[%d]: canonical provider/model identity is required", index)
}
if _, duplicate := enabled[canonical]; duplicate {
return catalog, fmt.Errorf("deploy/pi/settings.json enabledModels[%d]: duplicate identity %q", index, canonical)
}
enabled[canonical] = struct{}{}
piProvider, custom := piModels.Providers[providerID]
if !custom {
provider := catalog.Providers[providerID]
if len(provider.Models) == 0 {
provider = config.ModelProvider{
Authentication: config.ModelAuthentication{Mode: "pi_auth"},
Session: &config.ModelSessionAdapter{Mode: "pi_builtin"},
Models: make(map[string]config.CatalogModel),
}
}
provider.Models[modelID] = config.CatalogModel{Session: &config.SessionModel{}}
catalog.Providers[providerID] = provider
continue
}
model, found := findPiModel(piProvider.Models, modelID)
if !found {
return catalog, fmt.Errorf("deploy/pi/settings.json enabledModels[%d]: %q is missing from deploy/pi/models.json", index, canonical)
}
provider, exists := catalog.Providers[providerID]
if !exists {
auth, authErr := migratePiAuthentication(providerID, piProvider)
if authErr != nil {
return catalog, authErr
}
provider = config.ModelProvider{
Endpoint: &config.ModelEndpoint{BaseURL: piProvider.BaseURL}, Authentication: auth,
Session: &config.ModelSessionAdapter{Mode: "openai_compatible"},
Models: make(map[string]config.CatalogModel),
}
}
provider.Models[modelID] = config.CatalogModel{
Label: model.Name,
Session: &config.SessionModel{
Reasoning: model.Reasoning, Input: model.Input, Cost: model.Cost,
ContextWindow: model.ContextWindow, MaxTokens: model.MaxTokens,
Compatibility: model.Compat,
},
}
catalog.Providers[providerID] = provider
}
if _, ok := enabled[request.SessionDefault]; !ok {
return catalog, fmt.Errorf("session default %q is not enabled by deploy/pi/settings.json", request.SessionDefault)
}
metadataIDs := make(map[string]string, len(legacy.MetadataGeneration.Models))
for index, old := range legacy.MetadataGeneration.Models {
canonical, mergeErr := mergeMetadataModel(&catalog, old)
if mergeErr != nil {
return catalog, fmt.Errorf("metadataGeneration.models[%d]: %w", index, mergeErr)
}
if old.ID == "" {
return catalog, fmt.Errorf("metadataGeneration.models[%d]: id is required", index)
}
if _, duplicate := metadataIDs[old.ID]; duplicate {
return catalog, fmt.Errorf("metadataGeneration.models[%d]: duplicate legacy id %q", index, old.ID)
}
metadataIDs[old.ID] = canonical
}
if len(metadataIDs) > 0 {
mapped, ok := metadataIDs[legacy.MetadataGeneration.Default]
if !ok {
return catalog, errors.New("metadataGeneration.default does not identify one legacy model")
}
if mapped != request.SessionDefault {
return catalog, errors.New("migration_required: legacy metadata default differs from the requested interaction model; align the legacy default explicitly before migration")
}
} else if legacy.MetadataGeneration.Default != "" {
return catalog, errors.New("metadataGeneration.default is set without models")
}
return catalog, nil
}
func mergeMetadataModel(catalog *config.ModelCatalog, old legacyMetadataModel) (string, error) {
if old.LiteLLM.Provider == "" || old.LiteLLM.Model == "" {
return "", errors.New("litellm.provider and litellm.model are required")
}
type match struct{ provider, model string }
matches := make([]match, 0, 1)
for providerID, provider := range catalog.Providers {
if old.LiteLLM.Endpoint != nil && !sameEndpoint(provider.Endpoint, old.LiteLLM.Endpoint) {
continue
}
for modelID, model := range provider.Models {
upstream := model.UpstreamModel
if upstream == "" {
upstream = modelID
}
if upstream == old.LiteLLM.Model && (old.LiteLLM.Endpoint != nil || providerID == old.LiteLLM.Provider) {
matches = append(matches, match{providerID, modelID})
}
}
}
if len(matches) > 1 {
return "", errors.New("model identity matches more than one Pi provider")
}
providerID, modelID := old.LiteLLM.Provider, old.LiteLLM.Model
if len(matches) == 1 {
providerID, modelID = matches[0].provider, matches[0].model
}
provider, exists := catalog.Providers[providerID]
desiredAuth := config.ModelAuthentication{Mode: "none"}
if old.APIKeyEnv != "" {
desiredAuth = config.ModelAuthentication{Mode: "secret_env", APIKeyEnv: old.APIKeyEnv}
}
if !exists {
provider = config.ModelProvider{
Endpoint: old.LiteLLM.Endpoint, Authentication: desiredAuth,
MetadataGeneration: &config.ModelMetadataAdapter{LiteLLMProvider: old.LiteLLM.Provider},
Models: make(map[string]config.CatalogModel),
}
} else {
if provider.Authentication != desiredAuth {
return "", fmt.Errorf("provider %q authentication conflicts between Pi and metadata generation", providerID)
}
if old.LiteLLM.Endpoint != nil && !sameEndpoint(provider.Endpoint, old.LiteLLM.Endpoint) {
return "", fmt.Errorf("provider %q endpoint conflicts between Pi and metadata generation", providerID)
}
if provider.MetadataGeneration != nil && provider.MetadataGeneration.LiteLLMProvider != old.LiteLLM.Provider {
return "", fmt.Errorf("provider %q LiteLLM adapter is ambiguous", providerID)
}
provider.MetadataGeneration = &config.ModelMetadataAdapter{LiteLLMProvider: old.LiteLLM.Provider}
}
model := provider.Models[modelID]
if model.MetadataGeneration != nil {
return "", fmt.Errorf("canonical identity %q is claimed by more than one legacy metadata model", providerID+"/"+modelID)
}
if model.Label == "" {
model.Label = old.Label
}
if modelID != old.LiteLLM.Model {
model.UpstreamModel = old.LiteLLM.Model
}
model.MetadataGeneration = &config.MetadataGenerationModel{DisableThinking: old.LiteLLM.DisableThinking}
provider.Models[modelID] = model
catalog.Providers[providerID] = provider
return providerID + "/" + modelID, nil
}
func migratePiAuthentication(providerID string, provider piProvider) (config.ModelAuthentication, error) {
if provider.BaseURL == "" || provider.API != "openai-completions" {
return config.ModelAuthentication{}, fmt.Errorf("deploy/pi/models.json provider %q: only explicit openai-completions endpoints can migrate", providerID)
}
if strings.HasPrefix(provider.APIKey, "$") && len(provider.APIKey) > 1 {
return config.ModelAuthentication{Mode: "secret_env", APIKeyEnv: provider.APIKey[1:]}, nil
}
if provider.APIKey == "local" {
return config.ModelAuthentication{Mode: "none"}, nil
}
return config.ModelAuthentication{}, fmt.Errorf("deploy/pi/models.json provider %q: API key cannot be reconciled without guessing", providerID)
}
func findPiModel(models []piModel, id string) (piModel, bool) {
var found piModel
count := 0
for _, model := range models {
if model.ID == id {
found, count = model, count+1
}
}
return found, count == 1
}
func splitCanonical(value string) (string, string, bool) {
parts := strings.Split(value, "/")
return first(parts), second(parts), len(parts) == 2 && parts[0] != "" && parts[1] != ""
}
func first(parts []string) string {
if len(parts) > 0 {
return parts[0]
}
return ""
}
func second(parts []string) string {
if len(parts) > 1 {
return parts[1]
}
return ""
}
func sameEndpoint(left, right *config.ModelEndpoint) bool {
if left == nil || right == nil {
return left == nil && right == nil
}
return left.BaseURL == right.BaseURL && left.APIVersion == right.APIVersion
}
func decodeYAML(path string, target any) error {
source, err := readBounded(path)
if err != nil {
return err
}
decoder := yaml.NewDecoder(bytes.NewReader(source))
decoder.KnownFields(true)
if err := decoder.Decode(target); err != nil {
return errors.New("invalid or unsupported YAML")
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
return errors.New("multiple YAML documents are not supported")
}
return nil
}
func decodeJSON(path string, target any) error {
source, err := readBounded(path)
if err != nil {
return err
}
decoder := json.NewDecoder(bytes.NewReader(source))
decoder.DisallowUnknownFields()
if err := decoder.Decode(target); err != nil {
return errors.New("invalid or unsupported JSON")
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
return errors.New("trailing JSON is not supported")
}
return nil
}
func readBounded(path string) ([]byte, error) {
info, err := os.Lstat(path)
if err != nil || !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 || info.Size() < 1 || info.Size() > maxMigrationInputBytes {
return nil, errors.New("source is unavailable or unsafe")
}
return os.ReadFile(path)
}
func parseEnvironment(path string) (map[string]string, error) {
source, err := readBounded(path)
if err != nil {
return nil, err
}
return dotenv.ParseWithLookup(bytes.NewReader(source), os.LookupEnv)
}
func publishCandidate(path string, contents []byte) error {
directory := filepath.Dir(path)
if err := os.MkdirAll(directory, 0o700); err != nil {
return errors.New("create migration output directory")
}
temporary, err := os.CreateTemp(directory, ".installation-v2-*")
if err != nil {
return errors.New("create installation migration candidate")
}
name := temporary.Name()
published := false
defer func() {
if !published {
_ = os.Remove(name)
}
}()
if err := temporary.Chmod(0o600); err == nil {
_, err = temporary.Write(contents)
}
if closeErr := temporary.Close(); err == nil {
err = closeErr
}
if err != nil {
return errors.New("write installation migration candidate")
}
if err := os.Link(name, path); err != nil {
if os.IsExist(err) {
return errors.New("installation migration output already exists")
}
return errors.New("publish installation migration candidate")
}
published = true
_ = os.Remove(name)
return nil
}
// ParseDimensions is shared by the host command without accepting floats or signs.
func ParseDimensions(value string) (int, error) {
dimensions, err := strconv.Atoi(value)
if err != nil || dimensions <= 0 {
return 0, errors.New("embedding dimensions must be a positive integer")
}
return dimensions, nil
}