493 lines
12 KiB
Go
493 lines
12 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"
|
|
toolsrvclient "ollie/toolsrv/client"
|
|
"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
|
|
|
|
// Agent state
|
|
agents []*agent.Agent
|
|
log *olog.Logger
|
|
|
|
// Mutable state protected by mu
|
|
mu sync.RWMutex
|
|
name string // friendly name
|
|
toolsConn *toolsrvclient.Conn // multiplexed RPC connection
|
|
paused bool
|
|
|
|
// Autosave
|
|
saveMu sync.Mutex
|
|
saveDirty bool
|
|
saveTimer *time.Timer
|
|
|
|
// Models cache (per-session, refreshed every 24h)
|
|
modelsMu sync.Mutex
|
|
modelsCache string
|
|
modelsCacheAt time.Time
|
|
}
|
|
|
|
var sweepTmpOnce sync.Once
|
|
|
|
func ollieTmpDir() string {
|
|
return filepath.Join(util.DataDir(), "tmp")
|
|
}
|
|
|
|
func sweepStaleTmpDirs() {
|
|
sweepTmpOnce.Do(func() {
|
|
base := ollieTmpDir()
|
|
os.RemoveAll(base) //nolint:errcheck
|
|
os.MkdirAll(base, 0700) //nolint:errcheck
|
|
})
|
|
}
|
|
|
|
// --- 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
|
|
}
|
|
|
|
// ToolsConn returns the tool server connection.
|
|
func (s *Session) ToolsConn() *toolsrvclient.Conn {
|
|
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() *toolsrvclient.Conn {
|
|
s.mu.RLock()
|
|
k := s.Keeper
|
|
s.mu.RUnlock()
|
|
if k == nil {
|
|
return nil
|
|
}
|
|
conn, err := k.Dial()
|
|
if err != nil {
|
|
return nil
|
|
}
|
|
return conn
|
|
}
|
|
|
|
// SetToolsConn sets the tool server connection.
|
|
func (s *Session) SetToolsConn(conn *toolsrvclient.Conn) {
|
|
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() *toolsrvclient.Conn)
|
|
|
|
// 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 {
|
|
select {
|
|
case <-s.Ctx.Done():
|
|
return
|
|
case <-time.After(1 * time.Second):
|
|
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
|
|
}
|
|
|
|
// 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.
|
|
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) {
|
|
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.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
if s.paused {
|
|
return fmt.Errorf("session already paused")
|
|
}
|
|
if s.Cancel != nil {
|
|
s.Cancel()
|
|
}
|
|
if s.Keeper != nil {
|
|
s.Keeper.Close()
|
|
} else if s.Proc != nil {
|
|
s.Proc.Close()
|
|
}
|
|
if s.toolsConn != nil {
|
|
s.toolsConn.Close()
|
|
s.toolsConn = nil
|
|
}
|
|
s.paused = true
|
|
go s.saveSession()
|
|
PublishEvent("session."+s.ID+".pause", "")
|
|
return nil
|
|
}
|
|
|
|
// Resume restarts the tool server after a Pause.
|
|
func (s *Session) Resume() error {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
if !s.paused {
|
|
return fmt.Errorf("session not paused")
|
|
}
|
|
|
|
ctx, cancel := context.WithCancel(serverCtx)
|
|
s.Ctx = ctx
|
|
s.Cancel = cancel
|
|
|
|
// Ensure we always have a Keeper. If the session was restored from disk
|
|
// without infra, spawn a new tool server now.
|
|
if s.Keeper == nil {
|
|
cwd := ""
|
|
if len(s.agents) > 0 {
|
|
cwd = s.agents[0].Cwd()
|
|
}
|
|
if cwd == "" {
|
|
cwd, _ = os.Getwd()
|
|
}
|
|
infra, err := SetupToolServer(ToolServerConfig{
|
|
Ctx: ctx,
|
|
CWD: cwd,
|
|
RemoteTarget: s.Remote,
|
|
SessionID: s.ID,
|
|
Yolo: pkgYolo,
|
|
})
|
|
if err != nil {
|
|
return fmt.Errorf("resume: %w", err)
|
|
}
|
|
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)
|
|
}
|
|
} else {
|
|
// Keeper exists — just re-dial.
|
|
s.Keeper.SetContext(ctx)
|
|
conn, err := s.Keeper.Dial()
|
|
if err != nil {
|
|
return fmt.Errorf("resume failed: %w", err)
|
|
}
|
|
s.toolsConn = conn
|
|
// Update agent connections with fresh dials
|
|
for _, ag := range s.agents {
|
|
agConn, err := s.Keeper.Dial()
|
|
if err != nil {
|
|
return fmt.Errorf("resume agent dial: %w", err)
|
|
}
|
|
ag.SetToolServer(func() *toolsrvclient.Conn {
|
|
c, _ := s.Keeper.Dial()
|
|
return c
|
|
}, agConn)
|
|
ag.SetSessionEnv(s.ID)
|
|
}
|
|
}
|
|
|
|
s.paused = false
|
|
|
|
// 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", "")
|
|
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 ---
|
|
|
|
// 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
|
|
}
|
|
conn := s.ToolsConn()
|
|
if conn == nil {
|
|
return nil
|
|
}
|
|
// Ensure session ID is set on toolsrv before loading - the server uses this
|
|
// to scope the tool registry. Required because conn may have been reconnected.
|
|
conn.SetEnv("OLLIE_SESSION_ID", s.ID)
|
|
return LoadToolOnConn(conn, name)
|
|
}
|
|
|
|
// LoadToolOnConn loads a tool on the given connection.
|
|
func LoadToolOnConn(conn *toolsrvclient.Conn, 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()
|
|
}
|