352 lines
8.3 KiB
Go
352 lines
8.3 KiB
Go
package mgr
|
|
|
|
import (
|
|
"ollie/session"
|
|
"ollie/agent"
|
|
"context"
|
|
"fmt"
|
|
"ollie/backend"
|
|
"ollie/toolsrv"
|
|
"ollie/tools"
|
|
"olliesrv/prompts"
|
|
"os"
|
|
"path/filepath"
|
|
"sort"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
)
|
|
|
|
// --- Session Persistence ---
|
|
|
|
func (s *Manager) activeSessionsDir() string {
|
|
return filepath.Join(s.cfg.SessionsDir, "active")
|
|
}
|
|
|
|
func (s *Manager) persistSession(id string) {
|
|
s.mu.RLock()
|
|
sess, ok := s.sessions[id]
|
|
s.mu.RUnlock()
|
|
if !ok {
|
|
return
|
|
}
|
|
dir := s.activeSessionsDir()
|
|
os.MkdirAll(dir, 0700)
|
|
path := filepath.Join(dir, id+".json")
|
|
if err := sess.Core.SaveSession(path); err != nil {
|
|
s.cfg.Log.Error("persist session %s: %v", id, err)
|
|
}
|
|
}
|
|
|
|
func (s *Manager) removePersistedAgent(id string) {
|
|
path := filepath.Join(s.activeSessionsDir(), id+".json")
|
|
os.Remove(path)
|
|
}
|
|
|
|
func (s *Manager) saveAllSessions() {
|
|
s.mu.RLock()
|
|
ids := make([]string, 0, len(s.sessions))
|
|
for id := range s.sessions {
|
|
ids = append(ids, id)
|
|
}
|
|
s.mu.RUnlock()
|
|
for _, id := range ids {
|
|
s.persistSession(id)
|
|
}
|
|
}
|
|
|
|
func (s *Manager) restoreAllSessions() {
|
|
dir := s.activeSessionsDir()
|
|
entries, err := os.ReadDir(dir)
|
|
if err != nil {
|
|
return
|
|
}
|
|
|
|
// Load all persisted session JSONs (fast, sequential disk reads)
|
|
type loadedSession struct {
|
|
ps *agent.PersistedAgent
|
|
name string
|
|
}
|
|
var loaded []loadedSession
|
|
for _, e := range entries {
|
|
if !strings.HasSuffix(e.Name(), ".json") {
|
|
continue
|
|
}
|
|
path := filepath.Join(dir, e.Name())
|
|
ps, err := agent.LoadPersistedAgent(path)
|
|
if err != nil {
|
|
s.cfg.Log.Error("restore session %s: %v", e.Name(), err)
|
|
continue
|
|
}
|
|
loaded = append(loaded, loadedSession{ps: ps, name: e.Name()})
|
|
}
|
|
|
|
if len(loaded) == 0 {
|
|
return
|
|
}
|
|
|
|
// Restore sessions in parallel
|
|
var wg sync.WaitGroup
|
|
for _, ls := range loaded {
|
|
wg.Add(1)
|
|
go func(ls loadedSession) {
|
|
defer wg.Done()
|
|
if err := s.restoreSession(ls.ps); err != nil {
|
|
s.cfg.Log.Error("restore session %s: %v", ls.ps.ID, err)
|
|
}
|
|
}(ls)
|
|
}
|
|
wg.Wait()
|
|
}
|
|
|
|
func (s *Manager) restoreSession(ps *agent.PersistedAgent) error {
|
|
cwd := ps.CWD
|
|
if cwd == "" {
|
|
cwd, _ = os.Getwd()
|
|
}
|
|
agentName := ps.Agent
|
|
if agentName == "" {
|
|
agentName = "default"
|
|
}
|
|
sessID := ps.ID
|
|
|
|
cfg := LoadAgentConfig(s.cfg.AgentsDir, agentName, nil)
|
|
|
|
backendName := ps.Backend
|
|
if backendName == "" && cfg != nil && cfg.Backend != "" {
|
|
backendName = cfg.Backend
|
|
}
|
|
be, err := backend.NewWithName(backendName)
|
|
if err != nil {
|
|
return fmt.Errorf("backend: %w", err)
|
|
}
|
|
modelName := ps.Model
|
|
if modelName == "" && cfg != nil && cfg.Model != "" {
|
|
modelName = cfg.Model
|
|
}
|
|
if modelName == "" {
|
|
modelName = os.Getenv("OLLIE_MODEL")
|
|
}
|
|
if modelName != "" {
|
|
be.SetModel(modelName)
|
|
}
|
|
|
|
uname := s.nextUname()
|
|
var newToolServer func() toolsrv.Runner
|
|
var promptEnv []string
|
|
remoteTarget := ps.Remote
|
|
|
|
if remoteTarget != "" {
|
|
rsrv, dialErr := toolsrv.RemoteDial(s.cfg.Ctx, toolsrv.RemoteConfig{
|
|
SSHTarget: remoteTarget,
|
|
CWD: cwd,
|
|
})
|
|
if dialErr != nil {
|
|
return fmt.Errorf("remote dial: %w", dialErr)
|
|
}
|
|
newToolServer = func() toolsrv.Runner { return rsrv }
|
|
promptEnv = []string{
|
|
"PRIME_CWD=" + cwd,
|
|
"PRIME_PLATFORM=" + rsrv.Info.Platform,
|
|
"PRIME_IS_GIT_REPO=" + fmt.Sprintf("%v", rsrv.Info.IsGitRepo),
|
|
}
|
|
} else {
|
|
var execOpts []toolsrv.Option
|
|
if !s.cfg.NoMount {
|
|
}
|
|
if s.cfg.Strict {
|
|
execOpts = append(execOpts, toolsrv.WithStrict())
|
|
}
|
|
if s.cfg.Yolo {
|
|
execOpts = append(execOpts, toolsrv.WithYolo())
|
|
}
|
|
if s.cfg.ToolRegistry != nil {
|
|
execOpts = append(execOpts, toolsrv.WithToolRegistry(s.cfg.ToolRegistry, sessID))
|
|
}
|
|
if s.cfg.SkillsRegistry != nil {
|
|
execOpts = append(execOpts, toolsrv.WithSkillsRegistry(s.cfg.SkillsRegistry))
|
|
execOpts = append(execOpts, toolsrv.WithBuiltins(tools.Builtins()))
|
|
}
|
|
newToolServer = func() toolsrv.Runner {
|
|
srv := toolsrv.New(cwd)
|
|
for _, o := range execOpts {
|
|
o(srv)
|
|
}
|
|
return srv
|
|
}
|
|
promptEnv = agent.PromptEnv(cwd)
|
|
}
|
|
|
|
env := []string{"OLLIE_SESSION_ID=" + sessID, "OLLIE_UNAME=" + uname}
|
|
env = append(env, promptEnv...)
|
|
|
|
// Compute base layers (same as new-session path).
|
|
var spOverride string
|
|
if cfg != nil {
|
|
spOverride = cfg.SystemPrompt
|
|
}
|
|
sysPrompt := prompts.ResolveSystemPrompt(spOverride)
|
|
envMap := make(map[string]string)
|
|
for _, e := range env {
|
|
if k, v, ok := strings.Cut(e, "="); ok {
|
|
envMap[k] = v
|
|
}
|
|
}
|
|
opModel := prompts.OperationalModel(s.cfg.Enable9P, s.cfg.EnableDBus, envMap)
|
|
platform := "linux"
|
|
isGitRepo := false
|
|
for _, e := range promptEnv {
|
|
if k, v, ok := strings.Cut(e, "="); ok {
|
|
switch k {
|
|
case "PRIME_PLATFORM":
|
|
platform = v
|
|
case "PRIME_IS_GIT_REPO":
|
|
isGitRepo = v == "true"
|
|
}
|
|
}
|
|
}
|
|
envBlock := prompts.Environment(cwd, platform, isGitRepo, "")
|
|
|
|
toolSrv := newToolServer()
|
|
rt := agent.BuildRuntime(cfg, toolSrv, cwd, env, sysPrompt, opModel, envBlock)
|
|
|
|
restoredSession := agent.RestoreHistory(ps)
|
|
|
|
var sessPtr *Session
|
|
core := session.New(session.Config{
|
|
Backend: be,
|
|
AgentName: agentName,
|
|
AgentsDir: s.cfg.AgentsDir,
|
|
SessionsDir: s.cfg.SessionsDir,
|
|
SessionID: sessID,
|
|
AgentID: uname,
|
|
CWD: cwd,
|
|
Remote: remoteTarget,
|
|
History: restoredSession,
|
|
Runtime: rt,
|
|
NewToolServer: newToolServer,
|
|
PromptEnvExtra: promptEnv,
|
|
BaseLayers: []string{sysPrompt, opModel, envBlock},
|
|
Log: s.cfg.Sink.NewLogger("core"),
|
|
ReadPlanStep: func() string {
|
|
if sessPtr == nil {
|
|
return ""
|
|
}
|
|
sessPtr.mu.RLock()
|
|
data := make([]byte, len(sessPtr.plan))
|
|
copy(data, sessPtr.plan)
|
|
sessPtr.mu.RUnlock()
|
|
return session.NextUncheckedStep(data)
|
|
},
|
|
})
|
|
|
|
sessionCtx, sessionCancel := context.WithCancel(s.cfg.Ctx)
|
|
sess := NewSession(sessID, core, sessionCtx, sessionCancel)
|
|
sessPtr = sess
|
|
sess.uname = uname
|
|
sess.remote = remoteTarget
|
|
|
|
// Replay tail of persisted messages into the chat log so the GUI
|
|
// and `chat` file show recent history on restore.
|
|
replayMessagesToLog(sess, ps.Messages)
|
|
|
|
s.mu.Lock()
|
|
s.sessions[sessID] = sess
|
|
s.mu.Unlock()
|
|
|
|
s.cfg.Log.Info("restored session %s (backend=%s model=%s agent=%s)", sessID, backendName, modelName, agentName)
|
|
if s.cfg.OnSessionCreated != nil {
|
|
s.cfg.OnSessionCreated(sessID, sess)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (s *Manager) Shutdown() {
|
|
// Interrupt all in-progress turns and wait for them to finish
|
|
// before persisting state, so we capture the latest messages.
|
|
s.InterruptAll()
|
|
s.waitIdle(100*time.Millisecond, 5*time.Second)
|
|
s.saveAllSessions()
|
|
s.mu.Lock()
|
|
ids := make([]string, 0, len(s.sessions))
|
|
for id := range s.sessions {
|
|
ids = append(ids, id)
|
|
}
|
|
s.mu.Unlock()
|
|
for _, id := range ids {
|
|
s.mu.Lock()
|
|
sess := s.sessions[id]
|
|
delete(s.sessions, id)
|
|
s.mu.Unlock()
|
|
if sess != nil {
|
|
sess.Cancel()
|
|
sess.Core.Close()
|
|
|
|
s.cfg.Log.Info("shutdown session %s", id)
|
|
}
|
|
}
|
|
}
|
|
|
|
// waitIdle polls all sessions until they are idle or the timeout expires.
|
|
func (s *Manager) waitIdle(poll, timeout time.Duration) {
|
|
deadline := time.Now().Add(timeout)
|
|
for time.Now().Before(deadline) {
|
|
allIdle := true
|
|
s.mu.RLock()
|
|
for _, sess := range s.sessions {
|
|
if sess.Core.Agent().State() != "idle" {
|
|
allIdle = false
|
|
break
|
|
}
|
|
}
|
|
s.mu.RUnlock()
|
|
if allIdle {
|
|
return
|
|
}
|
|
time.Sleep(poll)
|
|
}
|
|
s.cfg.Log.Warn("shutdown: timed out waiting for sessions to become idle, saving anyway")
|
|
}
|
|
|
|
func (s *Manager) KillSession(id string) {
|
|
s.mu.Lock()
|
|
sess := s.sessions[id]
|
|
delete(s.sessions, id)
|
|
s.mu.Unlock()
|
|
if sess != nil {
|
|
sess.Cancel() // signal: context cancellation propagates to all agents
|
|
sess.Core.Close()
|
|
s.removePersistedAgent(id)
|
|
s.cfg.Log.Info("killed session %s", id)
|
|
if s.cfg.OnSessionKilled != nil {
|
|
s.cfg.OnSessionKilled(id)
|
|
}
|
|
}
|
|
}
|
|
|
|
func (s *Manager) index() []byte {
|
|
var sb strings.Builder
|
|
s.mu.RLock()
|
|
ids := make([]string, 0, len(s.sessions))
|
|
for id := range s.sessions {
|
|
ids = append(ids, id)
|
|
}
|
|
sort.Strings(ids)
|
|
for _, id := range ids {
|
|
sess := s.sessions[id]
|
|
sess.mu.RLock()
|
|
state := sess.Core.Agent().State()
|
|
cwd := sess.Core.CWD()
|
|
be := sess.Core.Agent().BackendName()
|
|
model := sess.Core.Agent().ModelName()
|
|
agent := sess.Core.Agent().Name()
|
|
sess.mu.RUnlock()
|
|
fmt.Fprintf(&sb, "%s\t%s\t%s\t%s\t%s\t%s\n", id, state, cwd, be, model, agent)
|
|
}
|
|
s.mu.RUnlock()
|
|
return []byte(sb.String())
|
|
}
|
|
|
|
// CreateSession creates a new agent session from key=value args.
|
|
// Returns the session ID on success.
|