package llama import ( "context" "encoding/json" "errors" "fmt" "io" "log/slog" "net/http" "os/exec" "strings" "sync" "syscall" "time" "gorm.io/gorm" "github.com/llamalink/llamalink/internal/config" "github.com/llamalink/llamalink/internal/db" ) var ( ErrModelNotFound = errors.New("model not found in registry") ErrModelDisabled = errors.New("model is disabled") ErrSwapInProgress = errors.New("model swap already in progress") ErrAlreadyLoaded = errors.New("model already loaded") ErrServerNotRunning = errors.New("llama-server not running") ) type Manager struct { cfg *config.Config db *gorm.DB mu sync.RWMutex state State proc *exec.Cmd done chan struct{} url string } func NewManager(cfg *config.Config, db *gorm.DB) *Manager { return &Manager{ cfg: cfg, db: db, done: make(chan struct{}), url: cfg.LlamaServerURL(), } } func (m *Manager) GetDB() *gorm.DB { return m.db } func (m *Manager) Start() error { m.mu.Lock() defer m.mu.Unlock() slog.Info("llama manager starting", "url", m.url) // Load default model on startup var model db.Model if err := m.db.Where("is_default = ? AND is_enabled = ?", true, true).First(&model).Error; err == nil { slog.Info("loading default model", "name", model.Name) if err := m.loadModelInternal(&model); err != nil { slog.Warn("failed to load default model", "error", err) m.state.Status = StatusFailed m.state.LastError = err.Error() return nil } } // Start health check loop go m.healthCheckLoop() return nil } func (m *Manager) Stop() { slog.Info("llama manager stopping") close(m.done) m.mu.Lock() defer m.mu.Unlock() if m.proc != nil && m.proc.Process != nil { ctx, cancel := context.WithTimeout(context.Background(), m.cfg.LlamaServerStopTimeoutDuration()) defer cancel() m.proc.SysProcAttr = &syscall.SysProcAttr{ Setpgid: true, } pgid, err := syscall.Getpgid(m.proc.Process.Pid) if err == nil { _ = syscall.Kill(-pgid, syscall.SIGTERM) } else { _ = m.proc.Process.Signal(syscall.SIGTERM) } <-ctx.Done() if m.proc.ProcessState == nil { _ = syscall.Kill(-pgid, syscall.SIGKILL) } } m.state = State{Status: StatusStopped} slog.Info("llama manager stopped") } func (m *Manager) IsReady() bool { m.mu.RLock() defer m.mu.RUnlock() return m.state.Status == StatusReady && m.proc != nil && m.proc.ProcessState != nil && !m.proc.ProcessState.Exited() } func (m *Manager) Status() Status { m.mu.RLock() defer m.mu.RUnlock() return m.state.Status } func (m *Manager) CurrentModel() string { m.mu.RLock() defer m.mu.RUnlock() return m.state.CurrentModel } func (m *Manager) GetStatus() *State { m.mu.RLock() defer m.mu.RUnlock() return &m.state } func (m *Manager) LoadModel(name string) error { m.mu.Lock() defer m.mu.Unlock() // Find model in DB var model db.Model if err := m.db.Where("name = ? AND is_enabled = ?", name, true).First(&model).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return ErrModelNotFound } return err } // Check current state if m.state.Status == StatusLoading || m.state.Status == StatusSwapping { if m.state.CurrentModel == name { return nil // Already loading this model } return fmt.Errorf("swap in progress for %s, try again later", m.state.TargetModel) } if m.state.CurrentModel == name && m.state.Status == StatusReady { return nil // Already loaded } return m.loadModelInternal(&model) } func (m *Manager) loadModelInternal(model *db.Model) error { isSwap := m.state.Status == StatusReady && m.state.CurrentModel != "" m.state.Status = StatusSwapping if !isSwap { m.state.Status = StatusLoading } m.state.TargetModel = model.Name m.state.LastError = "" now := time.Now() m.state.SwapStartedAt = &now slog.Info("loading model", "name", model.Name, "is_swap", isSwap) // Kill existing process if m.proc != nil && m.proc.Process != nil { m.terminateProcess() } // Build command cmd := m.buildCommand(model) m.proc = cmd if err := cmd.Start(); err != nil { m.state.Status = StatusFailed m.state.LastError = err.Error() return fmt.Errorf("failed to start llama-server: %w", err) } slog.Info("llama-server started", "pid", cmd.Process.Pid) m.state.PID = cmd.Process.Pid // Wait for server to be ready if err := m.waitUntilReady(); err != nil { m.state.Status = StatusFailed m.state.LastError = err.Error() return fmt.Errorf("model failed to start: %w", err) } m.state.CurrentModel = model.Name m.state.Status = StatusReady m.state.TargetModel = "" m.state.LoadedAt = &now // Update DB m.db.Model(model).Updates(map[string]interface{}{ "is_active": true, "loaded_at": now, }) slog.Info("model loaded successfully", "name", model.Name) return nil } func (m *Manager) buildCommand(model *db.Model) *exec.Cmd { args := []string{ "--model", model.ModelPath, "--alias", model.Alias, "--host", m.cfg.LlamaServerHost, "--port", fmt.Sprintf("%d", m.cfg.LlamaServerPort), "--ctx-size", fmt.Sprintf("%d", model.CtxSize), "--n-gpu-layers", fmt.Sprintf("%d", model.NGPULayers), } // Add extra args from JSON extraArgs := ParseModelExtraArgs(model.ExtraArgs) for k, v := range extraArgs { if bv, ok := v.(bool); ok && bv { args = append(args, "--"+k) } else if v != nil { args = append(args, "--"+k, fmt.Sprintf("%v", v)) } } cmd := exec.Command(m.cfg.LlamaServerBin, args...) cmd.Stdout = io.Discard cmd.Stderr = io.Discard // Set process group for clean kill cmd.SysProcAttr = &syscall.SysProcAttr{ Setpgid: true, } return cmd } func (m *Manager) terminateProcess() { if m.proc == nil || m.proc.Process == nil { return } slog.Info("terminating llama-server", "pid", m.proc.Process.Pid) ctx, cancel := context.WithTimeout(context.Background(), m.cfg.LlamaServerStopTimeoutDuration()) defer cancel() pgid, err := syscall.Getpgid(m.proc.Process.Pid) if err == nil { _ = syscall.Kill(-pgid, syscall.SIGTERM) } else { _ = m.proc.Process.Signal(syscall.SIGTERM) } done := make(chan error, 1) go func() { done <- m.proc.Wait() }() select { case <-ctx.Done(): if pgid, err := syscall.Getpgid(m.proc.Process.Pid); err == nil { _ = syscall.Kill(-pgid, syscall.SIGKILL) } case <-done: } m.proc = nil } func (m *Manager) waitUntilReady() error { ctx, cancel := context.WithTimeout(context.Background(), m.cfg.LlamaServerStartupTimeoutDuration()) defer cancel() ticker := time.NewTicker(500 * time.Millisecond) defer ticker.Stop() for { select { case <-ctx.Done(): return ctx.Err() case <-ticker.C: if m.checkHealth() { return nil } } } } func (m *Manager) checkHealth() bool { ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() req, err := http.NewRequestWithContext(ctx, "GET", m.url+"/health", nil) if err != nil { return false } resp, err := http.DefaultClient.Do(req) if err != nil { return false } defer resp.Body.Close() return resp.StatusCode == http.StatusOK } func (m *Manager) healthCheckLoop() { ticker := time.NewTicker(5 * time.Second) defer ticker.Stop() for { select { case <-m.done: return case <-ticker.C: m.healthCheck() } } } func (m *Manager) healthCheck() { m.mu.RLock() running := m.proc != nil && m.proc.Process != nil && m.proc.ProcessState != nil && !m.proc.ProcessState.Exited() m.mu.RUnlock() if !running && m.state.Status == StatusReady { m.mu.Lock() m.state.Status = StatusFailed m.state.LastError = "llama-server process died unexpectedly" m.mu.Unlock() slog.Error("llama-server process died", "current_model", m.state.CurrentModel) } } func (m *Manager) GetUsageStats() (totalRequests, totalTokens int64, avgLatencyMs float64) { now := time.Now() monthStart := time.Date(now.Year(), now.Month(), 1, 0, 0, 0, 0, time.UTC) var result struct { TotalRequests int64 TotalTokens int64 AvgLatency float64 } m.db.Model(&db.UsageLog{}). Where("created_at >= ?", monthStart). Select("COUNT(*) as total_requests, COALESCE(SUM(total_tokens), 0) as total_tokens, COALESCE(AVG(latency_ms), 0) as avg_latency"). Scan(&result) return result.TotalRequests, result.TotalTokens, result.AvgLatency } // ProxyRequest sends a request to the llama-server proxy func (m *Manager) ProxyRequest(ctx context.Context, method, path string, body io.Reader, headers map[string]string) (*http.Response, error) { if !m.IsReady() { return nil, ErrServerNotRunning } url := m.url + path req, err := http.NewRequestWithContext(ctx, method, url, body) if err != nil { return nil, err } req.Header.Set("Content-Type", "application/json") for k, v := range headers { req.Header.Set(k, v) } return http.DefaultClient.Do(req) } func ParseModelExtraArgs(extraArgs db.StringArray) map[string]interface{} { if len(extraArgs) == 0 { return nil } // If it's a JSON string, parse it if len(extraArgs) == 1 { var result map[string]interface{} if json.Unmarshal([]byte(extraArgs[0]), &result) == nil { return result } } // Otherwise assume key=value pairs result := make(map[string]interface{}) for _, arg := range extraArgs { parts := strings.SplitN(arg, "=", 2) if len(parts) == 2 { result[parts[0]] = parts[1] } else { result[arg] = true } } return result }