ollie/cmd/olliesrv/internal/session/session.go

611 lines
14 KiB
Go

package session
import (
"context"
"fmt"
"os"
"path/filepath"
"strings"
"sync"
"time"
"ollie/cmd/olliesrv/internal/agent"
"ollie/cmd/olliesrv/internal/backend"
"ollie/cmd/olliesrv/internal/toolclient"
olog "ollie/log"
"ollie/toolsrv/protocol"
"ollie/util"
)
// Session owns all session-level state: identity, agents, tool server lifecycle.
type Session struct {
// Identity (immutable after creation)
ID string // UUID
Uname string // user principal (numeric UID)
// Context lifecycle
Ctx context.Context
Cancel context.CancelFunc
// Tool server lifecycle (set during setup/teardown, not raced)
Proc *toolclient.Process
Keeper *toolclient.ProcessKeeper
Remote string
Yolo bool
// Agent state
agents []*agent.Agent
log *olog.Logger
// Mutable state protected by mu
mu sync.RWMutex
name string // friendly name
cwd string // session working directory
toolsConn *toolclient.ToolsrvConn // multiplexed RPC connection
paused bool
// Goal — session-level objective. Writing triggers a workflow.
goalMu sync.RWMutex
goalText string
goalStatus string // "", "running", "complete", "blocked", "error"
goalSignalMu sync.Mutex
goalSignalCh chan struct{}
workflow string // workflow profile name (default: "conductor")
variant string // workflow variant (empty = "default")
// Autosave
saveMu sync.Mutex
saveDirty bool
saveTimer *time.Timer
// Models cache (per-session, refreshed every 24h)
modelsMu sync.Mutex
modelsCache string
modelsCacheAt time.Time
}
func ollieTmpDir() string {
return filepath.Join(util.DataDir(), "tmp")
}
// --- Mutable field accessors (only for fields that race) ---
// Name returns the mutable friendly name.
func (s *Session) Name() string {
s.mu.RLock()
defer s.mu.RUnlock()
return s.name
}
// SetName sets the friendly name.
func (s *Session) SetName(name string) {
s.mu.Lock()
defer s.mu.Unlock()
s.name = name
}
// Cwd returns the session working directory.
func (s *Session) Cwd() string {
s.mu.RLock()
defer s.mu.RUnlock()
return s.cwd
}
// SetCwd sets the session working directory.
// Tilde and environment variables in cwd are expanded before storing.
func (s *Session) SetCwd(cwd string) {
s.mu.Lock()
s.cwd = util.ExpandHome(os.ExpandEnv(cwd))
s.mu.Unlock()
go s.saveSession()
}
// ToolsConn returns the tool server connection.
func (s *Session) ToolsConn() *toolclient.ToolsrvConn {
s.mu.RLock()
defer s.mu.RUnlock()
return s.toolsConn
}
// DialToolServer dials a fresh toolsrv connection via the keeper.
// Caller is responsible for closing the returned conn.
func (s *Session) DialToolServer() *toolclient.ToolsrvConn {
s.mu.RLock()
if s.paused {
s.mu.RUnlock()
return nil
}
k := s.Keeper
s.mu.RUnlock()
if k == nil {
s.log.Debug("DialToolServer: no keeper")
return nil
}
s.log.Debug("DialToolServer: dialing keeper paused=%v", s.IsPaused())
conn, err := k.Dial()
if err != nil {
s.log.Debug("DialToolServer: dial failed: %v", err)
return nil
}
s.log.Debug("DialToolServer: dial succeeded")
return conn
}
// SetToolsConn sets the tool server connection.
func (s *Session) SetToolsConn(conn *toolclient.ToolsrvConn) {
s.mu.Lock()
defer s.mu.Unlock()
s.toolsConn = conn
}
// BypassNotifyFunc is called when a bypass request arrives.
// It receives the request, session ID, and a dial function to get a connection for resolution.
type BypassNotifyFunc func(req *protocol.BypassRequest, sessionID string, dialFn func() *toolclient.ToolsrvConn)
// StartBypassLoop starts a goroutine that reads bypass requests from toolsrv
// and calls notifyFn for each one. This should be called after SetToolsConn.
func (s *Session) StartBypassLoop(notifyFn BypassNotifyFunc) {
go s.runBypassLoop(notifyFn)
}
func (s *Session) runBypassLoop(notifyFn BypassNotifyFunc) {
for {
// Get a fresh connection for each request (blocking reads don't multiplex well)
conn := s.DialToolServer()
if conn == nil {
timer := time.NewTimer(1 * time.Second)
select {
case <-s.Ctx.Done():
if !timer.Stop() {
<-timer.C
}
return
case <-timer.C:
continue
}
}
req, err := conn.ReadBypassPending(s.Ctx)
conn.Close()
if err != nil {
select {
case <-s.Ctx.Done():
return
default:
s.log.Debug("bypass read error: %v", err)
time.Sleep(100 * time.Millisecond)
continue
}
}
// Notify (non-blocking) - the notifier will write resolution when user responds
if notifyFn != nil {
notifyFn((*protocol.BypassRequest)(req), s.ID, s.DialToolServer)
}
}
}
// IsPaused returns true if the session is paused.
func (s *Session) IsPaused() bool {
s.mu.RLock()
defer s.mu.RUnlock()
return s.paused
}
// Goal returns the current goal text and status.
func (s *Session) Goal() (text, status string) {
s.goalMu.RLock()
defer s.goalMu.RUnlock()
return s.goalText, s.goalStatus
}
// SetGoal sets the goal text (the objective).
func (s *Session) SetGoal(text string) {
s.goalMu.Lock()
s.goalText = text
s.goalMu.Unlock()
s.notifyGoal()
go s.saveSession()
}
// SetGoalStatus updates the status field.
func (s *Session) SetGoalStatus(status string) {
s.goalMu.Lock()
s.goalStatus = status
s.goalMu.Unlock()
s.notifyGoal()
go s.saveSession()
}
// ClearGoal resets the entire goal state.
func (s *Session) ClearGoal() {
s.goalMu.Lock()
s.goalText = ""
s.goalStatus = ""
s.goalMu.Unlock()
s.notifyGoal()
go s.saveSession()
}
func (s *Session) notifyGoal() {
s.goalSignalMu.Lock()
if s.goalSignalCh != nil {
close(s.goalSignalCh)
}
s.goalSignalCh = make(chan struct{})
s.goalSignalMu.Unlock()
}
// GoalSignal returns a channel closed when the goal state changes.
func (s *Session) GoalSignal() <-chan struct{} {
s.goalSignalMu.Lock()
if s.goalSignalCh == nil {
s.goalSignalCh = make(chan struct{})
}
ch := s.goalSignalCh
s.goalSignalMu.Unlock()
return ch
}
// Workflow returns the workflow profile name (default: "none").
func (s *Session) Workflow() string {
s.mu.RLock()
w := s.workflow
s.mu.RUnlock()
return w
}
// SetWorkflow sets the workflow profile name.
func (s *Session) SetWorkflow(w string) {
s.mu.Lock()
s.workflow = w
s.mu.Unlock()
}
// Variant returns the workflow variant name (empty means "default").
func (s *Session) Variant() string {
s.mu.RLock()
v := s.variant
s.mu.RUnlock()
return v
}
// SetVariant sets the workflow variant name.
func (s *Session) SetVariant(v string) {
s.mu.Lock()
s.variant = v
s.mu.Unlock()
}
// FindAgent returns the agent matching the given name or ID, or nil.
func (s *Session) FindAgent(nameOrID string) *agent.Agent {
for _, ag := range s.agents {
if ag.Name() == nameOrID || ag.ID() == nameOrID {
return ag
}
}
return nil
}
// Agents returns the full agent slice.
func (s *Session) Agents() []*agent.Agent { return s.agents }
// AddAgent appends an agent to the session.
func (s *Session) AddAgent(ag *agent.Agent) {
s.agents = append(s.agents, ag)
if s.Ctx != nil {
go agent.ConsumeFeed(s.Ctx, ag)
}
}
// RemoveAgent removes an agent by ID and closes it.
// Any peer relationships to this agent are cleaned up from remaining agents.
func (s *Session) RemoveAgent(id string) bool {
idx := -1
var removed *agent.Agent
for i, ag := range s.agents {
if ag.ID() == id {
idx = i
removed = ag
break
}
}
if idx < 0 {
return false
}
s.agents[idx].Close()
s.agents = append(s.agents[:idx], s.agents[idx+1:]...)
// Remove peer links pointing to the removed agent.
name := removed.Name()
for _, ag := range s.agents {
ag.RemovePeer(name)
}
return true
}
// Close releases resources for this session.
func (s *Session) Close() {
s.log.Debug("Close() session=%q", s.ID)
if s.Cancel != nil {
s.Cancel()
}
s.flushSave()
for _, ag := range s.agents {
ag.Close()
}
if s.ID != "" {
os.RemoveAll(filepath.Join(ollieTmpDir(), s.ID)) //nolint:errcheck
}
}
// SetEnv stores a session-scoped variable and propagates to all agents.
func (s *Session) SetEnv(key, value string) {
for _, ag := range s.agents {
ag.SetEnv(key, value)
}
}
// SetSessionID renames the session.
func (s *Session) SetSessionID(newID string) error {
oldID := s.ID
if oldID == newID {
return nil
}
if pkgSessionsDir != "" && oldID != "" {
for _, suffix := range []string{".json", ".compaction.jsonl"} {
oldPath := s.activeSessionPath(oldID, suffix)
if _, err := os.Stat(oldPath); err == nil {
if err := os.Rename(oldPath, s.activeSessionPath(newID, suffix)); err != nil {
return fmt.Errorf("rename %s: %w", suffix, err)
}
}
}
}
s.ID = newID
oldTemp := filepath.Join(ollieTmpDir(), oldID)
newTemp := filepath.Join(ollieTmpDir(), newID)
if _, err := os.Stat(oldTemp); err == nil {
os.Rename(oldTemp, newTemp) //nolint:errcheck
}
for _, ag := range s.agents {
ag.SetSessionEnv(newID)
}
return nil
}
// Pause stops the tool server to save resources.
func (s *Session) Pause() error {
s.log.Debug("Pause: begin")
s.mu.Lock()
if s.paused {
s.mu.Unlock()
return fmt.Errorf("session already paused")
}
s.paused = true
cancel := s.Cancel
keeper := s.Keeper
proc := s.Proc
conn := s.toolsConn
s.Cancel = nil
s.Keeper = nil
s.Proc = nil
s.toolsConn = nil
s.mu.Unlock()
if cancel != nil {
cancel()
}
if keeper != nil {
if err := keeper.Close(); err != nil {
s.log.Debug("Pause: closing keeper: %v", err)
}
} else if proc != nil {
if err := proc.Close(); err != nil {
s.log.Debug("Pause: closing process: %v", err)
}
}
if conn != nil {
conn.Close()
}
s.log.Debug("Pause: complete")
go s.saveSession()
PublishEvent("session."+s.ID+".pause", "")
return nil
}
// Resume restarts the tool server after a Pause.
func (s *Session) Resume() error {
s.log.Debug("Resume: begin")
s.mu.Lock()
if !s.paused {
s.mu.Unlock()
return fmt.Errorf("session not paused")
}
ctx, cancel := context.WithCancel(serverCtx)
s.Ctx = ctx
s.Cancel = cancel
s.log.Debug("Resume: context created")
// Start fresh toolsrv infrastructure.
s.log.Info("Resume: starting toolsrv")
cwd := s.cwd
if cwd == "" {
cwd, _ = os.Getwd()
}
if s.Remote != "" {
s.log.Info("Resume: spawning remote toolsrv target=%s cwd=%s", s.Remote, cwd)
} else {
s.log.Info("Resume: spawning local toolsrv cwd=%s", cwd)
}
infra, err := SetupToolServer(ToolServerConfig{
Ctx: ctx,
CWD: cwd,
RemoteTarget: s.Remote,
SessionID: s.ID,
Yolo: s.Yolo,
})
if err != nil {
s.log.Error("Resume: toolsrv startup failed: %v", err)
cancel()
s.mu.Unlock()
return fmt.Errorf("resume: %w", err)
}
s.log.Info("Resume: toolsrv started")
s.Proc = infra.Proc
s.Keeper = infra.Keeper
s.toolsConn = infra.ToolsConn
for _, ag := range s.agents {
ag.SetToolServer(infra.NewToolServer, infra.NewToolServer())
ag.SetSessionEnv(s.ID)
}
s.paused = false
s.mu.Unlock()
s.log.Debug("Resume: marked active")
// Start bypass approval loop for the new session context.
if pkgBypassNotify != nil {
s.StartBypassLoop(pkgBypassNotify)
}
// Rebuild full runtimes for restored agents (they have stub runtimes from persist).
for _, ag := range s.agents {
if err := rebuildAgentRuntime(ag, s.ID); err != nil {
pkgLog.Error("resume agent %s: %v", ag.ID(), err)
}
}
// Start feed consumers for restored agents (context now available).
for _, ag := range s.agents {
go agent.ConsumeFeed(ctx, ag)
}
go s.saveSession()
PublishEvent("session."+s.ID+".resume", "")
s.log.Debug("Resume: complete")
return nil
}
// rebuildAgentRuntime rebuilds the full runtime for an agent (used on resume).
func rebuildAgentRuntime(ag *agent.Agent, sessID string) error {
cfg := agent.LoadConfig(pkgAgentsDir, ag.Profile(), os.Open)
layers := BuildPromptLayers(cfg, ag.Cwd(), sessID, ag.ID(), "linux", false, "")
env := []string{"OLLIE_SESSION_ID=" + sessID, "OLLIE_UNAME=" + ag.ID()}
rt := agent.BuildRuntime(cfg, ag.ToolServer(), ag.Cwd(), env, layers.SystemPrompt, layers.EnvBlock)
be, err := backend.NewWithName(cfg.Backend)
if err != nil {
return fmt.Errorf("backend: %w", err)
}
if cfg.Model != "" {
be.SetModel(cfg.Model)
}
rt.Backend = be
ag.SetRuntime(rt)
LoadAutoLoadTools(cfg, ag.ToolServer(), sessID, ag.ID(), func(f string, a ...any) {
pkgLog.Error("session %s agent %s: "+f, append([]any{sessID, ag.ID()}, a...)...)
})
return nil
}
// --- Tool loading ---
// LoadToolOnConn loads a tool on the given connection.
func LoadToolOnConn(conn *toolclient.ToolsrvConn, name string) error {
if conn == nil {
return nil
}
return conn.LoadTool(name)
}
// --- Autosave ---
func (s *Session) activeSessionPath(id, suffix string) string {
return filepath.Join(pkgSessionsDir, "active", id+suffix)
}
func (s *Session) saveSession() {
s.saveMu.Lock()
s.saveDirty = true
if s.saveTimer == nil {
s.saveTimer = time.AfterFunc(2*time.Second, s.flushSave)
}
s.saveMu.Unlock()
}
func (s *Session) flushSave() {
s.saveMu.Lock()
dirty := s.saveDirty
s.saveDirty = false
if s.saveTimer != nil {
s.saveTimer.Stop()
s.saveTimer = nil
}
s.saveMu.Unlock()
if !dirty || s.ID == "" || pkgSessionsDir == "" {
return
}
if err := PersistSession(s.Name()); err != nil {
s.log.Error("session save: %v", err)
}
}
// Save writes the current session state via PersistSession.
func (s *Session) Save() error {
return PersistSession(s.Name())
}
// --- Constructor for empty session ---
// NewEmpty creates an empty session without an agent.
func NewEmpty(id string, ctx context.Context, cancel context.CancelFunc) *Session {
name := id
if idx := strings.IndexByte(id, '-'); idx > 0 {
name = id[:idx]
}
log := pkgSink.NewLogger("session")
return &Session{
ID: id,
name: name,
Ctx: ctx,
Cancel: cancel,
agents: make([]*agent.Agent, 0),
log: log,
}
}
// --- Models cache ---
const modelsCacheTTL = 24 * time.Hour
// CachedListModels returns a cached model list for this session.
func (s *Session) CachedListModels() string {
s.modelsMu.Lock()
if s.modelsCache != "" && time.Since(s.modelsCacheAt) < modelsCacheTTL {
result := s.modelsCache
s.modelsMu.Unlock()
return result
}
s.modelsMu.Unlock()
agents := s.Agents()
if len(agents) == 0 {
return ""
}
models := agents[0].ListModels()
result := strings.Join(models, "\n")
s.modelsMu.Lock()
s.modelsCache = result
s.modelsCacheAt = time.Now()
s.modelsMu.Unlock()
return result
}
// InvalidateModelsCache clears the cached model list.
func (s *Session) InvalidateModelsCache() {
s.modelsMu.Lock()
s.modelsCache = ""
s.modelsCacheAt = time.Time{}
s.modelsMu.Unlock()
}