307 lines
7.2 KiB
Go
307 lines
7.2 KiB
Go
package session
|
|
|
|
import (
|
|
"context"
|
|
"crypto/rand"
|
|
"fmt"
|
|
"io"
|
|
"os"
|
|
"path/filepath"
|
|
"strconv"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/simonfxr/pubsub"
|
|
"ollie/agent"
|
|
"ollie/backend"
|
|
olog "ollie/log"
|
|
"ollie/paths"
|
|
"ollie/tools"
|
|
)
|
|
|
|
// Config is the configuration for creating a session.
|
|
type Config struct {
|
|
Backend backend.Backend
|
|
ModelName string
|
|
AgentName string
|
|
AgentsDir string
|
|
SessionsDir string
|
|
SessionID string
|
|
AgentID string
|
|
CWD string
|
|
History *agent.History
|
|
Runtime *agent.Runtime
|
|
NewToolServer func() tools.Runner
|
|
NewBackend func(string) (backend.Backend, error)
|
|
Log *olog.Logger
|
|
MaxSteps int
|
|
ReadPlanStep func() string
|
|
ListHandlers map[string]func() []string
|
|
PromptEnvExtra []string
|
|
Remote string
|
|
BaseLayers []string
|
|
}
|
|
|
|
// Session is the concrete session type. It owns session-level state and
|
|
// delegates agent operations to its owned Agent.
|
|
type Session struct {
|
|
id string
|
|
bus *pubsub.Bus
|
|
envMu sync.RWMutex
|
|
env map[string]string
|
|
plan []byte
|
|
prevPrompt string
|
|
|
|
r *agent.Agent
|
|
log *olog.Logger
|
|
sessionsDir string
|
|
listHandlers map[string]func() []string
|
|
remote string
|
|
mu sync.RWMutex
|
|
auditLog *olog.Logger
|
|
|
|
saveMu sync.Mutex
|
|
saveDirty bool
|
|
saveTimer *time.Timer
|
|
}
|
|
|
|
// NewSessionID generates a unique, lexicographically sortable session identifier.
|
|
func NewSessionID() string {
|
|
b := make([]byte, 3)
|
|
rand.Read(b) //nolint:errcheck
|
|
return strconv.FormatInt(time.Now().UnixNano(), 10) + "-" + fmt.Sprintf("%06x", b)
|
|
}
|
|
|
|
// 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 {
|
|
if p := os.Getenv("OLLIE_TMP_PATH"); p != "" {
|
|
return p
|
|
}
|
|
return filepath.Join(os.TempDir(), "ollie")
|
|
}
|
|
|
|
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)
|
|
}
|
|
|
|
bus := pubsub.NewBus()
|
|
auditLog := log.Sub("audit")
|
|
|
|
a := &Session{
|
|
id: cfg.SessionID,
|
|
bus: bus,
|
|
env: make(map[string]string),
|
|
log: log,
|
|
auditLog: auditLog,
|
|
sessionsDir: cfg.SessionsDir,
|
|
remote: cfg.Remote,
|
|
listHandlers: cfg.ListHandlers,
|
|
}
|
|
|
|
a.r = agent.NewAgent(agent.AgentCfg{
|
|
History: cfg.History,
|
|
Runtime: rt,
|
|
AgentName: cfg.AgentName,
|
|
AgentsDir: cfg.AgentsDir,
|
|
AgentID: cfg.AgentID,
|
|
CWD: paths.ExpandHome(cfg.CWD),
|
|
BaseLayers: cfg.BaseLayers,
|
|
PromptEnvExtra: cfg.PromptEnvExtra,
|
|
NewToolServer: cfg.NewToolServer,
|
|
NewBackend: cfg.NewBackend,
|
|
Bus: bus,
|
|
Log: log,
|
|
AuditLog: auditLog,
|
|
SessionID: cfg.SessionID,
|
|
StartupMsgs: rt.Messages,
|
|
ReadPlanStep: cfg.ReadPlanStep,
|
|
SaveSession: a.saveSession,
|
|
FlushSave: a.flushSave,
|
|
})
|
|
|
|
a.r.SetSessionEnv(a.id)
|
|
return a
|
|
}
|
|
|
|
// Close releases resources for this session.
|
|
func (a *Session) Close() {
|
|
a.log.Debug("Close() session=%q", a.id)
|
|
a.flushSave()
|
|
a.r.Close()
|
|
if a.id != "" {
|
|
os.RemoveAll(filepath.Join(ollieTmpDir(), a.id)) //nolint:errcheck
|
|
}
|
|
}
|
|
|
|
// SetEnv stores a session-scoped variable and propagates it to the agent.
|
|
func (a *Session) SetEnv(key, value string) {
|
|
a.envMu.Lock()
|
|
a.env[key] = value
|
|
a.envMu.Unlock()
|
|
a.r.SetEnv(key, value)
|
|
}
|
|
|
|
func (a *Session) Agent() *agent.Agent { return a.r }
|
|
func (a *Session) Bus() *pubsub.Bus { return a.bus }
|
|
|
|
// CWD returns the current working directory for tool execution.
|
|
func (a *Session) CWD() string {
|
|
if c := a.r.Cwd(); c != "" {
|
|
return c
|
|
}
|
|
wd, _ := os.Getwd()
|
|
return wd
|
|
}
|
|
|
|
// SetCWD validates and sets the working directory.
|
|
func (a *Session) SetCWD(dir string) error {
|
|
dir = paths.ExpandHome(dir)
|
|
if dir != "" {
|
|
if _, err := os.Stat(dir); err != nil {
|
|
return fmt.Errorf("cwd: %w", err)
|
|
}
|
|
}
|
|
a.r.SetCWD(dir)
|
|
return nil
|
|
}
|
|
|
|
// SetSessionID renames the session.
|
|
func (a *Session) SetSessionID(newID string) error {
|
|
oldID := a.id
|
|
if oldID == newID {
|
|
return nil
|
|
}
|
|
if a.sessionsDir != "" && oldID != "" {
|
|
for _, suffix := range []string{".json", ".compaction.jsonl"} {
|
|
oldPath := a.activeSessionPath(oldID, suffix)
|
|
if _, err := os.Stat(oldPath); err == nil {
|
|
if err := os.Rename(oldPath, a.activeSessionPath(newID, suffix)); err != nil {
|
|
return fmt.Errorf("rename %s: %w", suffix, err)
|
|
}
|
|
}
|
|
}
|
|
}
|
|
a.id = newID
|
|
a.r.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
|
|
}
|
|
a.r.SetSessionEnv(newID)
|
|
return nil
|
|
}
|
|
|
|
// WaitChange blocks until the named field changes from current.
|
|
func (a *Session) WaitChange(ctx context.Context, field, current string) (string, bool) {
|
|
if field == agent.WatchState {
|
|
return a.r.WaitChange(ctx, field, current)
|
|
}
|
|
// Other fields — read via agent methods, use agent's change signal.
|
|
read := func() string {
|
|
switch field {
|
|
case "usage":
|
|
return a.r.UsageStr()
|
|
case "ctxsz":
|
|
return a.r.CtxSz()
|
|
case "cwd":
|
|
return a.CWD()
|
|
case "agent":
|
|
return a.r.Name()
|
|
}
|
|
return ""
|
|
}
|
|
stop := context.AfterFunc(ctx, func() { a.r.BroadcastChange() })
|
|
defer stop()
|
|
for ctx.Err() == nil {
|
|
if v := read(); v != current {
|
|
return v, true
|
|
}
|
|
a.r.WaitForChange(ctx)
|
|
}
|
|
return "", false
|
|
}
|
|
|
|
func (a *Session) activeSessionPath(id, suffix string) string {
|
|
return filepath.Join(a.sessionsDir, "active", id+suffix)
|
|
}
|
|
|
|
func (a *Session) saveSession() {
|
|
a.saveMu.Lock()
|
|
a.saveDirty = true
|
|
if a.saveTimer == nil {
|
|
a.saveTimer = time.AfterFunc(2*time.Second, a.flushSave)
|
|
}
|
|
a.saveMu.Unlock()
|
|
}
|
|
|
|
func (a *Session) flushSave() {
|
|
a.saveMu.Lock()
|
|
dirty := a.saveDirty
|
|
a.saveDirty = false
|
|
if a.saveTimer != nil {
|
|
a.saveTimer.Stop()
|
|
a.saveTimer = nil
|
|
}
|
|
a.saveMu.Unlock()
|
|
if !dirty || a.id == "" || a.sessionsDir == "" {
|
|
return
|
|
}
|
|
path := a.activeSessionPath(a.id, ".json")
|
|
if err := os.MkdirAll(filepath.Dir(path), 0700); err != nil {
|
|
a.log.Error("session save: %v", err)
|
|
return
|
|
}
|
|
if err := a.r.SaveFull(path, a.id, a.CWD(), a.remote); err != nil {
|
|
a.log.Error("session save: %v", err)
|
|
}
|
|
}
|
|
|
|
// SaveSession writes the current session state to the given path.
|
|
func (a *Session) SaveSession(path string) error {
|
|
return a.r.SaveFull(path, a.id, a.CWD(), a.remote)
|
|
}
|