ollie/session/session.go

633 lines
15 KiB
Go

package session
import (
"context"
"crypto/rand"
"encoding/json"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"sync"
"time"
"ollie/agent"
"ollie/backend"
olog "ollie/log"
"ollie/paths"
"ollie/toolsrv"
)
// Config is the configuration for creating a session with an agent.
type Config struct {
Backend backend.Backend
ModelName string
Profile string // config profile name (e.g. "default")
AgentsDir string
SessionsDir string
SessionID string
AgentID string
CWD string
History *agent.History
Runtime *agent.Runtime
NewToolServer func() *toolsrv.Conn
NewBackend func(string) (backend.Backend, error)
Log *olog.Logger
MaxSteps int
ReadPlanStep func() string
ListHandlers map[string]func() []string
PromptEnvExtra []string
Remote string
SystemPrompt string
EnvBlock string
Output agent.EventHandler
}
// Session owns all session-level state: identity, agents, tool server lifecycle.
type Session struct {
// Identity
id string // immutable UUID, set at creation
name string // mutable friendly name, defaults to id prefix
uname string // immutable user principal (numeric UID)
// Context lifecycle
ctx context.Context
cancel context.CancelFunc
// Agent state
agents []*agent.Agent
log *olog.Logger
auditLog *olog.Logger
sessionsDir string
listHandlers map[string]func() []string
remote string
// Environment
envMu sync.RWMutex
env map[string]string
plan []byte
prevPrompt string
// Tool server lifecycle
proc *toolsrv.Process // toolsrv subprocess
keeper *toolsrv.ProcessKeeper // resilient process manager
toolsConn *toolsrv.Conn // multiplexed RPC connection
paused bool // true when paused
// Tool restrictions
disallowTools map[string]struct{}
// Synchronization
mu sync.RWMutex
// Autosave
saveMu sync.Mutex
saveDirty bool
saveTimer *time.Timer
}
// NewSessionID generates a UUIDv4 session identifier.
func NewSessionID() string {
b := make([]byte, 16)
rand.Read(b) //nolint:errcheck
b[6] = (b[6] & 0x0f) | 0x40 // version 4
b[8] = (b[8] & 0x3f) | 0x80 // variant 10xx
return fmt.Sprintf("%08x-%04x-%04x-%04x-%012x",
b[0:4], b[4:6], b[6:8], b[8:10], b[10:16])
}
// NextUncheckedStep returns the first unchecked step from plan bytes.
func NextUncheckedStep(data []byte) string {
for _, line := range strings.Split(string(data), "\n") {
trimmed := strings.TrimSpace(line)
if strings.HasPrefix(trimmed, "- [ ]") {
return strings.TrimSpace(trimmed[5:])
}
}
return ""
}
var sweepTmpOnce sync.Once
func ollieTmpDir() string {
return filepath.Join(paths.DataDir(), "tmp")
}
func sweepStaleTmpDirs() {
sweepTmpOnce.Do(func() {
base := ollieTmpDir()
os.RemoveAll(base) //nolint:errcheck
os.MkdirAll(base, 0700) //nolint:errcheck
})
}
// New creates a session with an owned agent from the given configuration.
func New(cfg Config) *Session {
sweepStaleTmpDirs()
if cfg.ModelName != "" {
cfg.Backend.SetModel(cfg.ModelName)
}
if cfg.NewBackend == nil {
cfg.NewBackend = backend.NewWithName
}
rt := cfg.Runtime
if rt == nil {
rt = &agent.Runtime{}
}
rt.Backend = cfg.Backend
if cfg.MaxSteps > 0 {
rt.MaxSteps = cfg.MaxSteps
}
if cfg.SessionID != "" {
os.MkdirAll(filepath.Join(ollieTmpDir(), cfg.SessionID), 0700) //nolint:errcheck
}
log := cfg.Log
if log == nil {
log = olog.NewWriter("core", olog.LevelError+1, io.Discard, io.Discard)
}
auditLog := log.Sub("audit")
name := cfg.SessionID
if idx := strings.IndexByte(cfg.SessionID, '-'); idx > 0 {
name = cfg.SessionID[:idx]
}
s := &Session{
id: cfg.SessionID,
name: name,
env: make(map[string]string),
agents: make([]*agent.Agent, 0, 1),
log: log,
auditLog: auditLog,
sessionsDir: cfg.SessionsDir,
remote: cfg.Remote,
listHandlers: cfg.ListHandlers,
}
ag := agent.NewAgent(agent.AgentCfg{
History: cfg.History,
Runtime: rt,
Profile: cfg.Profile,
AgentsDir: cfg.AgentsDir,
ID: cfg.AgentID,
Cwd: paths.ExpandHome(cfg.CWD),
SystemPrompt: cfg.SystemPrompt,
EnvBlock: cfg.EnvBlock,
PromptEnvExtra: cfg.PromptEnvExtra,
NewToolServer: cfg.NewToolServer,
NewBackend: cfg.NewBackend,
Output: cfg.Output,
Log: log,
AuditLog: auditLog,
SessionID: cfg.SessionID,
StartupMsgs: rt.Messages,
ReadPlanStep: cfg.ReadPlanStep,
Save: s.saveSession,
Flush: s.flushSave,
})
s.agents = append(s.agents, ag)
ag.SetSessionEnv(s.id)
return s
}
// 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 }
// AgentAt returns the agent at the given index, or nil.
func (s *Session) AgentAt(idx int) *agent.Agent {
if idx < 0 || idx >= len(s.agents) {
return nil
}
return s.agents[idx]
}
// AgentCount returns the number of agents.
func (s *Session) AgentCount() int { return len(s.agents) }
// SaveFuncs returns save and flush callbacks suitable for use as
// SaveFuncs returns the save and flush callbacks for use in
// agent.AgentCfg.Save and Flush when adding agents to this session.
func (s *Session) SaveFuncs() (save func(), flush func()) {
return s.saveSession, s.flushSave
}
// AddAgent appends an agent to the session.
func (s *Session) AddAgent(ag *agent.Agent) {
s.agents = append(s.agents, ag)
ag.SetSessionEnv(s.id)
}
// RemoveAgent removes an agent by ID and closes it.
// Returns false if the agent was not found.
func (s *Session) RemoveAgent(id string) bool {
idx := -1
for i, ag := range s.agents {
if ag.ID() == id {
idx = i
break
}
}
if idx < 0 {
return false
}
s.agents[idx].Close()
s.agents = append(s.agents[:idx], s.agents[idx+1:]...)
return true
}
// Close releases resources for this session.
func (s *Session) Close() {
s.log.Debug("Close() session=%q", s.id)
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) {
s.envMu.Lock()
s.env[key] = value
s.envMu.Unlock()
for _, ag := range s.agents {
ag.SetEnv(key, value)
}
}
// Remote returns the remote target (e.g., SSH host) if any.
func (s *Session) Remote() string {
return s.remote
}
// SetRemote sets the remote target for this session.
func (s *Session) SetRemote(remote string) {
s.remote = remote
}
// SetSessionID renames the session.
func (s *Session) SetSessionID(newID string) error {
oldID := s.id
if oldID == newID {
return nil
}
if s.sessionsDir != "" && 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
for _, ag := range s.agents {
ag.RenamePreamble(oldID, 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
}
func (s *Session) activeSessionPath(id, suffix string) string {
return filepath.Join(s.sessionsDir, "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 == "" || s.sessionsDir == "" {
return
}
// Use PersistSession for consistent format (writes Name field correctly)
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())
}
// --- Identity accessors ---
// ID returns the immutable session UUID.
func (s *Session) ID() string { return s.id }
// 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
}
// Uname returns the immutable user principal.
func (s *Session) Uname() string { return s.uname }
// SetUname sets the user principal (should only be called during init).
func (s *Session) SetUname(uname string) {
s.mu.Lock()
defer s.mu.Unlock()
s.uname = uname
}
// --- Context lifecycle ---
// Ctx returns the session context.
func (s *Session) Ctx() context.Context { return s.ctx }
// Cancel cancels the session context.
func (s *Session) Cancel() {
if s.cancel != nil {
s.cancel()
}
}
// SetContext sets the session context and cancel func.
func (s *Session) SetContext(ctx context.Context, cancel context.CancelFunc) {
s.ctx = ctx
s.cancel = cancel
}
// --- Tool server lifecycle ---
// Proc returns the tool server process.
func (s *Session) Proc() *toolsrv.Process {
s.mu.RLock()
defer s.mu.RUnlock()
return s.proc
}
// SetProc sets the tool server process.
func (s *Session) SetProc(proc *toolsrv.Process) {
s.mu.Lock()
defer s.mu.Unlock()
s.proc = proc
}
// Keeper returns the process keeper.
func (s *Session) Keeper() *toolsrv.ProcessKeeper {
s.mu.RLock()
defer s.mu.RUnlock()
return s.keeper
}
// SetKeeper sets the process keeper.
func (s *Session) SetKeeper(keeper *toolsrv.ProcessKeeper) {
s.mu.Lock()
defer s.mu.Unlock()
s.keeper = keeper
}
// ToolsConn returns the tool server connection.
func (s *Session) ToolsConn() *toolsrv.Conn {
s.mu.RLock()
defer s.mu.RUnlock()
return s.toolsConn
}
// SetToolsConn sets the tool server connection.
func (s *Session) SetToolsConn(conn *toolsrv.Conn) {
s.mu.Lock()
defer s.mu.Unlock()
s.toolsConn = conn
}
// IsPaused returns true if the session is paused.
func (s *Session) IsPaused() bool {
s.mu.RLock()
defer s.mu.RUnlock()
return s.paused
}
// SetPaused sets the paused state directly (used during restore).
func (s *Session) SetPaused(paused bool) {
s.mu.Lock()
defer s.mu.Unlock()
s.paused = paused
}
// IsConnected returns true if the tool server connection is alive.
// Returns false if the session is paused.
func (s *Session) IsConnected() bool {
s.mu.RLock()
paused := s.paused
conn := s.toolsConn
s.mu.RUnlock()
if paused || conn == nil {
return false
}
return conn.Ping() == nil
}
// Pause stops the tool server to save resources.
// It cancels the session context (stopping all agent operations),
// kills the tool server process, and marks the session as paused.
func (s *Session) Pause() error {
s.mu.Lock()
defer s.mu.Unlock()
if s.paused {
return fmt.Errorf("session already paused")
}
// Cancel the session context to stop all agent operations.
if s.cancel != nil {
s.cancel()
}
// Close the tool server process.
if s.keeper != nil {
s.keeper.Close()
} else if s.proc != nil {
s.proc.Close()
}
// Close any existing connection.
if s.toolsConn != nil {
s.toolsConn.Close()
s.toolsConn = nil
}
s.paused = true
// Persist the paused state.
go s.saveSession()
return nil
}
// Resume restarts the tool server after a Pause.
// It creates a fresh session context and respawns the tool server.
// For sessions restored in paused state (no keeper), it sets up the tool server from scratch.
func (s *Session) Resume() error {
s.mu.Lock()
defer s.mu.Unlock()
if !s.paused {
return fmt.Errorf("session not paused")
}
// Create a fresh context since the old one was cancelled on pause.
ctx, cancel := context.WithCancel(serverCtx)
s.ctx = ctx
s.cancel = cancel
// If we have a keeper, use it to respawn. Otherwise set up from scratch.
if s.keeper != nil {
s.keeper.SetContext(ctx)
conn, err := s.keeper.Dial()
if err != nil {
return fmt.Errorf("resume failed: %w", err)
}
s.toolsConn = conn
} else {
// Session was restored paused without infra - set up per-agent.
for i, ag := range s.agents {
cwd := ag.Cwd()
if cwd == "" {
cwd, _ = os.Getwd()
}
// First agent spawns the tool server, others reuse
var reuseFrom *InfraConfig
if s.proc != nil {
reuseFrom = &InfraConfig{
Proc: s.proc,
Keeper: s.keeper,
ToolsConn: s.toolsConn,
}
}
infra, err := SetupToolServer(ToolServerConfig{
Ctx: ctx,
CWD: cwd,
RemoteTarget: s.remote,
Yolo: pkgYolo,
ReuseFrom: reuseFrom,
})
if err != nil {
return fmt.Errorf("resume setup for agent %s failed: %w", ag.ID(), err)
}
// First agent sets session infra
if i == 0 {
s.proc = infra.Proc
s.keeper = infra.Keeper
s.toolsConn = infra.ToolsConn
}
ag.SetToolServer(infra.NewToolServer, infra.ToolsConn)
}
}
s.paused = false
// Persist the resumed state.
go s.saveSession()
return nil
}
// --- Tool loading ---
// SetDisallowTools sets the tool disallow list.
func (s *Session) SetDisallowTools(disallow map[string]struct{}) {
s.mu.Lock()
defer s.mu.Unlock()
s.disallowTools = disallow
}
// LoadTool loads a tool into the tool server.
func (s *Session) LoadTool(name string, ag *agent.Agent) error {
name = strings.TrimSpace(name)
if name == "" {
return nil
}
s.mu.RLock()
_, blocked := s.disallowTools[name]
s.mu.RUnlock()
if blocked {
return fmt.Errorf("tool %q is disallowed for this session", name)
}
var conn *toolsrv.Conn
if s.toolsConn != nil {
conn = s.toolsConn
} else if ag != nil {
conn = ag.ToolServer()
}
return LoadToolOnConn(conn, name)
}
// LoadToolOnConn loads a tool on the given connection.
func LoadToolOnConn(conn *toolsrv.Conn, name string) error {
if conn == nil {
return nil
}
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
args, _ := json.Marshal(map[string]string{"name": name})
if _, err := conn.CallTool(ctx, "tool_load", json.RawMessage(args)); err != nil {
return fmt.Errorf("remote load: %w", err)
}
return nil
}
// --- 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,
env: make(map[string]string),
agents: make([]*agent.Agent, 0),
log: log,
auditLog: log.Sub("audit"),
sessionsDir: pkgSessionsDir,
}
}