From d9920b42c72736de9cbfeb701dfa4aea1a69b980 Mon Sep 17 00:00:00 2001 From: Levi Neely Date: Wed, 29 Jul 2026 19:29:49 +0200 Subject: [PATCH] session: merge harness.go into session.go MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The Session struct is defined in session.go — named after what it is. No more 'harness' concept; that was an artifact of the interface+impl pattern we removed. --- session/harness.go | 1572 ------------------------------------ session/session.go | 1567 +++++++++++++++++++++++++++++++++++ session/waitchange_test.go | 2 +- 3 files changed, 1568 insertions(+), 1573 deletions(-) delete mode 100644 session/harness.go diff --git a/session/harness.go b/session/harness.go deleted file mode 100644 index bf09fcf..0000000 --- a/session/harness.go +++ /dev/null @@ -1,1572 +0,0 @@ -package session - -import ( - "context" - "crypto/rand" - "encoding/json" - "errors" - "fmt" - "io" - "os" - "path/filepath" - "runtime/debug" - "slices" - "strconv" - "strings" - "sync" - "sync/atomic" - "syscall" - "time" - - "github.com/simonfxr/pubsub" - "ollie/backend" - olog "ollie/log" - "ollie/paths" - "ollie/tools" -) - -// toolClassifier reports whether a named tool is safe to run concurrently -// with other read-class tools. nil means treat all tools as serial. -type toolClassifier func(name string) bool - -// BuildRuntime constructs a Runtime from a pre-configured Dispatcher and -// optional agent config. cwd sets the working directory reported in the -// system prompt; if empty, the process working directory is used. -// env provides additional environment variables injected into prompt resolution -// subprocesses (e.g. OLLIE_SESSION_ID=xxx). -// The caller is responsible for registering all servers on d before calling this. -func BuildRuntime(cfg *AgentConfig, d tools.Dispatcher, cwd string, env []string, baseLayers ...string) *Runtime { - var messages []string - - var allToolInfos []tools.ToolInfo - var allTools []backend.Tool - serverOf := make(map[string]string) - - if cfg == nil || cfg.ToolsEnabled() { - var listErr error - allToolInfos, listErr = d.ListTools() - if listErr != nil { - messages = append(messages, fmt.Sprintf("list tools: %v", listErr)) - } - for _, t := range allToolInfos { - serverOf[t.Name] = t.Server - } - // Only built-in executors (with InputSchema) become backend tools. - // Named tool scripts are promoted via the tool registry and appear in the preamble. - allTools = toolInfosToBackend(allToolInfos) - - // Append named tool scripts for preamble listing only. - allToolInfos = append(allToolInfos, tools.DiscoverTools()...) - } - - hooks := Hooks{} - var preamble string - var genParams backend.GenerationParams - var maxSteps int - if cfg != nil { - for k, v := range cfg.Hooks { - hooks[k] = []string(v) - } - if resolved, err := resolvePrompt(cfg.Prompt, cwd, env); err != nil { - fmt.Fprintf(os.Stderr, "resolve prompt: %v\n", err) - } else { - preamble = resolved - } - genParams = backend.GenerationParams{ - MaxTokens: cfg.MaxTokens, - MaxCompletionTokens: cfg.MaxCompletionTokens, - Temperature: cfg.Temperature, - TopP: cfg.TopP, - TopK: cfg.TopK, - MinP: cfg.MinP, - TopA: cfg.TopA, - FrequencyPenalty: cfg.FrequencyPenalty, - PresencePenalty: cfg.PresencePenalty, - RepetitionPenalty: cfg.RepetitionPenalty, - ThinkingBudget: cfg.Reasoning, - ReasoningEffort: cfg.ReasoningEffort, - IncludeReasoning: cfg.IncludeReasoning, - ResponseFormat: cfg.ResponseFormat, - Stop: cfg.Stop, - Verbosity: cfg.Verbosity, - } - maxSteps = cfg.MaxSteps - if len(cfg.AllowTools) > 0 { - if srv, ok := d.GetServer("execute"); ok { - if rs, ok := srv.(tools.ToolRestrictionSetter); ok { - rs.SetAllowTools(cfg.AllowTools) - } - } - } - } - - exec := func(ctx context.Context, name string, args json.RawMessage) (string, []backend.ContentBlock, error) { - // Dynamically resolve the server for this tool — handles - // lazy-promoted tools that appeared after BuildRuntime. - infos, listErr := d.ListTools() - if listErr != nil { - return "", nil, listErr - } - server := "" - for _, t := range infos { - if t.Name == name { - server = t.Server - break - } - } - if server == "" { - return "", nil, fmt.Errorf("unknown tool: %s", name) - } - raw, err := d.Dispatch(ctx, server, name, args) - if err != nil { - return "", nil, err - } - text, blocks, isErr := extractToolResult(raw) - if isErr { - return "", nil, fmt.Errorf("%s", text) - } - return text, blocks, nil - } - - var classify toolClassifier - if srv, ok := d.GetServer("execute"); ok { - if pc, ok := srv.(tools.ParallelClassifier); ok { - classify = pc.IsParallelRead - } - } - - var tierFn func(string, json.RawMessage) ResultTier - if srv, ok := d.GetServer("execute"); ok { - if tc, ok := srv.(tools.TierClassifier); ok { - tierFn = func(name string, args json.RawMessage) ResultTier { - switch tc.ResultTierArgs(name, args) { - case "cold": - return TierCold - case "warm": - return TierWarm - default: - return TierHot - } - } - } - } - - var backendName, modelName, compactionModel string - if cfg != nil { - backendName = cfg.Backend - modelName = cfg.Model - compactionModel = cfg.CompactionModel - } - - // Prepend base layers (system prompt, operational model, environment) - // before the agent-specific preamble. - if len(baseLayers) > 0 { - var prefix strings.Builder - for _, layer := range baseLayers { - if layer != "" { - prefix.WriteString(layer) - prefix.WriteByte('\n') - } - } - if prefix.Len() > 0 { - preamble = prefix.String() + preamble - } - } - - // Append compact tool surface listing (name + description). - // Full prompts are available on-demand via /tools write (net/dns pattern). - var toolListing strings.Builder - for _, ti := range allToolInfos { - if ti.Description != "" && ti.Server == "" { - // Only named tool scripts (Server==""), not built-in executors - fmt.Fprintf(&toolListing, "- **%s** — %s\n", ti.Name, ti.Description) - } - } - if toolListing.Len() > 0 { - preamble += "\n# Available Tools\n\n" + toolListing.String() - } - - return &Runtime{ - Dispatcher: d, - Tools: allTools, - Exec: exec, - ClassifyTool: classify, - ClassifyTier: tierFn, - Hooks: hooks, - Preamble: preamble, - GenParams: genParams, - MaxSteps: maxSteps, - CfgBackend: backendName, - CfgModel: modelName, - CompactionModel: compactionModel, - Messages: messages, - } -} - -// DefaultPromptsDir returns the default directory for prompt templates. -func DefaultPromptsDir() string { - return paths.CfgDir() + "/prompts" -} - -// PromptsDirs returns all prompt directories from OLLIE_PROMPTS_PATH (colon-separated). -func PromptsDirs() []string { - if p := os.Getenv("OLLIE_PROMPTS_PATH"); p != "" { - return strings.Split(p, ":") - } - return []string{DefaultPromptsDir()} -} - -// AgentsDirs returns all agent directories from OLLIE_AGENTS_PATH (colon-separated). -func AgentsDirs() []string { - if p := os.Getenv("OLLIE_AGENTS_PATH"); p != "" { - return strings.Split(p, ":") - } - return []string{paths.CfgDir() + "/agents"} -} - -// AgentConfigPath resolves the config file path for a named agent. -// It searches all directories in OLLIE_AGENTS_PATH, falling back to agentsDir. -func AgentConfigPath(agentsDir, name string) string { - for _, dir := range AgentsDirs() { - p := dir + "/" + name + ".json" - if _, err := os.Stat(p); err == nil { - return p - } - } - return agentsDir + "/" + name + ".json" -} - -// NewSessionID generates a unique, lexicographically sortable session identifier. -// Format: - — sortable by creation time, unique -// even if two sessions are created within the same nanosecond. -func NewSessionID() string { - b := make([]byte, 3) - rand.Read(b) //nolint:errcheck - return strconv.FormatInt(time.Now().UnixNano(), 10) + "-" + fmt.Sprintf("%06x", b) -} - -// NewResponseID generates a unique identifier for a single assistant response. -func NewResponseID() string { - b := make([]byte, 3) - rand.Read(b) //nolint:errcheck - return "resp_" + strconv.FormatInt(time.Now().UnixNano(), 10) + "_" + fmt.Sprintf("%06x", b) -} - -// NewReactionID generates a unique identifier for a user reaction. -func NewReactionID() string { - b := make([]byte, 3) - rand.Read(b) //nolint:errcheck - return "react_" + strconv.FormatInt(time.Now().UnixNano(), 10) + "_" + fmt.Sprintf("%06x", b) -} - -// infoEvent wraps a plain-text message as an info Event. -func infoEvent(text string) Event { - return Event{Role: "info", Content: text + "\n"} -} - -// actionHandle holds the cancel function for the current agent turn. -type actionHandle struct { - cancel context.CancelCauseFunc -} - -// Config is the configuration for creating an agent. -type Config struct { - Backend backend.Backend - ModelName string // if non-empty, overrides backend's default model - AgentName string - AgentsDir string - SessionsDir string - SessionID string - Uname string // immutable user principal for 9P identity - CWD string // working directory for tool execution and system prompt - History *History - Runtime *Runtime - NewDispatcher func() tools.Dispatcher - NewBackend func(string) (backend.Backend, error) // if nil, defaults to backend.NewWithName - Log *olog.Logger // if nil, logging is disabled - // MaxSteps overrides the agent JSON maxSteps when non-zero. - // 0 means use the value from the agent config (or unlimited if absent). - MaxSteps int - // ReadPlanStep, if non-nil, is called to get the next unchecked plan step. - // When provided, it is used directly instead of reading from the filesystem. - ReadPlanStep func() string - // ListHandlers provides harness-specific list implementations for slash - // commands such as /skills, /tools, /agents. Keyed by command name without - // the leading slash. Commands with no entry emit nothing. - ListHandlers map[string]func() []string - // PromptEnvExtra holds PRIME_* env vars for prompt resolution. - // For remote sessions these come from HostInfo; for local, from PromptEnv(). - // Stored so /agent reloads use the correct values. - PromptEnvExtra []string - // Remote is the SSH target for remote execution (empty = local). - Remote string - // BaseLayers are the system prompt, operational model, and environment - // block prepended to the agent preamble. Stored so /agent reloads - // preserve the server-injected context. - BaseLayers []string -} - -// harness is the Core implementation. It owns all harness and session state -// but has no knowledge of how output is rendered. -type Session struct { - // Session-level state - id string - uname string - cwd string - state string // "idle", "thinking", "calling: " - reply string - bus *pubsub.Bus - fifo Fifo - envMu sync.RWMutex - env map[string]string - plan []byte - prevPrompt string - peers map[string]bool - - // Agent - r *Agent - log *olog.Logger - sessionsDir string - readPlanStep func() string - listHandlers map[string]func() []string - turnError func(ctx context.Context, errType, errMsg string) HookResult // overridable for tests - remote string // SSH target for remote execution - startupMessages []string - toolCallCount atomic.Int64 - pendingInject atomic.Pointer[string] - mu sync.RWMutex - changeMu sync.Mutex - changeCond *sync.Cond - submitMu sync.Mutex // serializes Submit calls (commands + turns) - auditLog *olog.Logger - - // Debounced session persistence - saveMu sync.Mutex - saveDirty bool - saveTimer *time.Timer - -} - -// ToolCallCount returns the total number of tool calls executed in this -// session. The counter is monotonically increasing and never resets. -// Blocked calls (pre-tool hook exit 2) are not counted. -func (a *Session) ToolCallCount() int64 { - return a.toolCallCount.Load() -} - -// SetEnv stores a session-scoped variable and propagates it to the execute server. -func (a *Session) SetEnv(key, value string) { - a.envMu.Lock() - a.env[key] = value - a.envMu.Unlock() - if a.r.runtime == nil || a.r.runtime.Dispatcher == nil { - return - } - if srv, ok := a.r.runtime.Dispatcher.GetServer("execute"); ok { - if es, ok := srv.(tools.EnvSetter); ok { - es.SetEnv(key, value) - } - } -} - -// pushSessionEnv injects OLLIE_SESSION_ID into the execute server subprocess env. -func (a *Session) pushSessionEnv() { - if a.r.runtime == nil || a.r.runtime.Dispatcher == nil || a.id == "" { - return - } - if srv, ok := a.r.runtime.Dispatcher.GetServer("execute"); ok { - if es, ok := srv.(tools.EnvSetter); ok { - es.SetEnv("OLLIE_SESSION_ID", a.id) - if a.uname != "" { - es.SetEnv("OLLIE_UNAME", a.uname) - } - } - } -} - -// pushLockDir sets the flock directory on the execute server to the session tmpdir. -func (a *Session) pushLockDir() { - if a.r.runtime == nil || a.r.runtime.Dispatcher == nil || a.id == "" { - return - } - -} - - -var sweepTmpOnce sync.Once - -// ollieTmpDir returns the base temp directory for session tmpdirs. -func ollieTmpDir() string { - if p := os.Getenv("OLLIE_TMP_PATH"); p != "" { - return p - } - return filepath.Join(os.TempDir(), "ollie") -} - -// readNextUnchecked reads a plan file and returns the first unchecked step. -func readNextUnchecked(path string) string { - data, err := os.ReadFile(path) - if err != nil || len(data) == 0 { - return "" - } - return NextUncheckedStep(data) -} - -// 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 "" -} - -// sweepStaleTmpDirs removes session tmpdirs left by a previous crash. -// All session tmpdirs are owned by a single process, so anything -// present at startup is stale. -func sweepStaleTmpDirs() { - sweepTmpOnce.Do(func() { - base := ollieTmpDir() - os.RemoveAll(base) //nolint:errcheck - os.MkdirAll(base, 0700) //nolint:errcheck - }) -} - -// New creates an 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 = &Runtime{} - } - // Store the backend on the runtime so it's the single source of truth. - rt.Backend = cfg.Backend - // A non-zero MaxSteps in Config takes precedence over the - // value loaded from the agent JSON. - if cfg.MaxSteps > 0 { - rt.MaxSteps = cfg.MaxSteps - } - var readPlanStep func() string - if cfg.SessionID != "" { - os.MkdirAll(filepath.Join(ollieTmpDir(), cfg.SessionID), 0700) //nolint:errcheck - } - if cfg.ReadPlanStep != nil { - readPlanStep = cfg.ReadPlanStep - } - - log := cfg.Log - if log == nil { - log = olog.NewWriter("core", olog.LevelError+1, io.Discard, io.Discard) - } - - a := &Session{ - id: cfg.SessionID, - uname: cfg.Uname, - cwd: paths.ExpandHome(cfg.CWD), - state: "idle", - bus: pubsub.NewBus(), - env: make(map[string]string), - r: &Agent{ - history: cfg.History, - runtime: rt, - agentName: cfg.AgentName, - agentsDir: cfg.AgentsDir, - promptEnvExtra: cfg.PromptEnvExtra, - baseLayers: cfg.BaseLayers, - newDispatcher: cfg.NewDispatcher, - newBackend: cfg.NewBackend, - }, - log: log, - auditLog: log.Sub("audit"), - sessionsDir: cfg.SessionsDir, - remote: cfg.Remote, - startupMessages: rt.Messages, - readPlanStep: readPlanStep, - listHandlers: cfg.ListHandlers, - } - a.changeCond = sync.NewCond(&a.changeMu) - a.turnError = func(_ context.Context, errType, errMsg string) HookResult { - hookCtx, cancel := context.WithTimeout(context.Background(), time.Duration(hookTimeout)*time.Second) - defer cancel() - return a.r.runtime.Hooks.Run(hookCtx, HookTurnError, map[string]string{ - "session_id": a.id, - "cwd": a.CWD(), - "model": a.r.runtime.Backend.Model(), - "error_type": errType, - "error": errMsg, - }, a.log) - } - a.pushSessionEnv() - a.pushLockDir() - return a -} - -// Close releases resources for this session, including its tmpdir. -func (a *Session) Close() { - a.log.Debug("Close() session=%q", a.id) - a.flushSave() - if a.r.runtime != nil && a.r.runtime.Dispatcher != nil { - if srv, ok := a.r.runtime.Dispatcher.GetServer("execute"); ok { - if c, ok := srv.(interface{ Close() }); ok { - a.log.Debug("Close() calling execute.Close()") - c.Close() - } - } - } - if a.id != "" { - os.RemoveAll(filepath.Join(ollieTmpDir(), a.id)) //nolint:errcheck - } -} - -// execServer returns the execute server if available, or nil. -func (a *Session) execServer() interface{} { - if a.r.runtime == nil || a.r.runtime.Dispatcher == nil { - return nil - } - srv, _ := a.r.runtime.Dispatcher.GetServer("execute") - return srv -} - -func (a *Session) Detach() bool { - if srv := a.execServer(); srv != nil { - if d, ok := srv.(interface{ Detach() bool }); ok { - return d.Detach() - } - } - return false -} - -func (a *Session) ListDetached() []DetachedInfo { - if srv := a.execServer(); srv != nil { - type listDetacher interface { - ListDetachedRaw() []any - } - if ld, ok := srv.(listDetacher); ok { - raw := ld.ListDetachedRaw() - out := make([]DetachedInfo, 0, len(raw)) - for _, r := range raw { - if m, ok := r.(map[string]any); ok { - di := DetachedInfo{} - if v, ok := m["pid"].(int); ok { - di.PID = v - } - if v, ok := m["command"].(string); ok { - di.Command = v - } - if v, ok := m["started"].(int64); ok { - di.Started = v - } - if v, ok := m["exited"].(bool); ok { - di.Exited = v - } - if v, ok := m["exit_code"].(int); ok { - di.ExitCode = v - } - out = append(out, di) - } - } - return out - } - } - return nil -} - -func (a *Session) SignalDetached(pid, signal int) error { - if srv := a.execServer(); srv != nil { - type signaler interface { - SignalDetached(int, syscall.Signal) error - } - if sg, ok := srv.(signaler); ok { - return sg.SignalDetached(pid, syscall.Signal(signal)) - } - } - return fmt.Errorf("no execute server available") -} - -func (a *Session) GetDetachedOutput(pid int) (string, error) { - if srv := a.execServer(); srv != nil { - type outputGetter interface { - GetDetachedOutput(int) (string, error) - } - if og, ok := srv.(outputGetter); ok { - return og.GetDetachedOutput(pid) - } - } - return "", fmt.Errorf("no execute server available") -} - -func (a *Session) DismissDetached(pid int) bool { - if srv := a.execServer(); srv != nil { - type dismisser interface { - DismissDetached(int) bool - } - if d, ok := srv.(dismisser); ok { - return d.DismissDetached(pid) - } - } - return false -} - -func (a *Session) InjectSystemEvent(content string) { - a.Queue("\n" + content + "\n") -} - -// classifyReaction returns a category and description for a reaction emoji. -func classifyReaction(emoji string) (category, description string, positive bool) { - switch emoji { - case "👍", "✅": - return "positive", "The response was good. Keep doing what you're doing.", true - case "🚀", "🎉": - return "excellent", "The response was exactly what was wanted.", true - case "👎", "❌": - return "negative", "The response was wrong or unhelpful.", false - case "💩", "🤬": - return "terrible", "The response was fundamentally wrong. Stop this approach entirely and reassess from scratch.", false - case "🤔": - return "confused", "The response was unclear or confusing.", false - default: - return "unknown", "", false - } -} - -func (a *Session) Reactions() map[string]string { - result := make(map[string]string) - if a.r.history == nil { - return result - } - for _, reaction := range a.r.history.Reactions { - result[reaction.ResponseID] = reaction.Emoji - } - return result -} - -func (a *Session) React(emoji string) { - _ = a.ReactTo("", emoji) -} - -func (a *Session) ReactTo(responseID, emoji string) error { - if a.r.history == nil { - return fmt.Errorf("no active session") - } - category, _, _ := classifyReaction(emoji) - if category == "unknown" { - return fmt.Errorf("unsupported reaction: %s", emoji) - } - if responseID == "" { - for i := len(a.r.history.messages) - 1; i >= 0; i-- { - if a.r.history.messages[i].Role == "assistant" { - responseID = a.r.history.messages[i].ID - break - } - } - } - if responseID == "" { - return fmt.Errorf("no assistant response to react to") - } - found := false - for i := range a.r.history.messages { - if a.r.history.messages[i].Role == "assistant" && a.r.history.messages[i].ID == responseID { - found = true - break - } - } - if !found { - return fmt.Errorf("assistant response not found: %s", responseID) - } - - reaction := Reaction{ID: NewReactionID(), ResponseID: responseID, Emoji: emoji, Category: category, CreatedAt: time.Now()} - replaced := false - for i := range a.r.history.Reactions { - if a.r.history.Reactions[i].ResponseID == responseID { - if a.r.history.Reactions[i].Emoji == emoji { - return nil - } - a.r.history.Reactions[i] = reaction - replaced = true - break - } - } - if !replaced { - a.r.history.Reactions = append(a.r.history.Reactions, reaction) - } - a.r.history.recomputeReactionCounts() - a.saveSession() - return nil -} - -func (a *Session) AgentName() string { - v := a.r.agentName - a.log.Debug("AgentName() = %q", v) - return v -} -func (a *Session) BackendName() string { - v := a.r.runtime.Backend.Name() - a.log.Debug("BackendName() = %q", v) - return v -} -func (a *Session) ModelName() string { - v := a.r.runtime.Backend.Model() - a.log.Debug("ModelName() = %q", v) - return v -} - -func (a *Session) State() string { - return a.state -} - -func (a *Session) notifyChange() { - a.changeMu.Lock() - a.changeCond.Broadcast() - a.changeMu.Unlock() -} - -func (a *Session) setState(state string) { - a.state = state - a.log.Debug("state -> %q", state) - a.notifyChange() -} - -// WaitChange blocks until the named field changes from current, then returns -// the new value. Returns ("", false) if ctx is cancelled. -func (a *Session) WaitChange(ctx context.Context, field, current string) (string, bool) { - read := func() string { - switch field { - case WatchState: - return a.State() - case WatchUsage: - return a.Usage() - case WatchCtxSz: - return a.CtxSz() - case WatchCWD: - return a.CWD() - case WatchAgent: - return a.AgentName() - } - return "" - } - - // context.AfterFunc fires in a separate goroutine when ctx is done, - // broadcasting to unblock any waiters. - stop := context.AfterFunc(ctx, func() { - a.changeMu.Lock() - a.changeCond.Broadcast() - a.changeMu.Unlock() - }) - defer stop() - - a.changeMu.Lock() - defer a.changeMu.Unlock() - for ctx.Err() == nil { - if v := read(); v != current { - return v, true - } - a.changeCond.Wait() - } - return "", false -} - - -func (a *Session) Reply() string { - r := a.reply - a.log.Debug("Reply() len=%d", len(r)) - return r -} - -// CWD returns the current working directory for tool execution. -func (a *Session) CWD() string { - if a.cwd != "" { - a.log.Debug("CWD() = %q", a.cwd) - return a.cwd - } - wd, _ := os.Getwd() - a.log.Debug("CWD() = %q (from getwd)", wd) - return wd -} - -// SetCWD changes the working directory for tool execution and updates the -// system prompt. Returns an error if the path does not exist. -func (a *Session) SetCWD(dir string) error { - a.log.Debug("SetCWD(%q)", dir) - dir = paths.ExpandHome(dir) - if dir != "" { - if _, err := os.Stat(dir); err != nil { - return fmt.Errorf("cwd: %w", err) - } - } - oldCwd := a.cwd - a.cwd = dir - // Update cwd references in the system prompt. - if oldCwd != "" && dir != "" && oldCwd != dir { - a.r.runtime.Preamble = strings.ReplaceAll(a.r.runtime.Preamble, oldCwd, dir) - } - // Propagate to any tool server that knows how to handle it (e.g. execute). - if a.r.runtime != nil && a.r.runtime.Dispatcher != nil { - if srv, ok := a.r.runtime.Dispatcher.GetServer("execute"); ok { - if ws, ok := srv.(tools.CWDSetter); ok { - ws.SetCWD(dir) - } - } - } - a.notifyChange() - return nil -} - -// SetSessionID renames the session. It updates the in-memory ID, renames -// persisted files on disk, and propagates to the execute server env. -func (a *Session) SetSessionID(newID string) error { - a.log.Debug("SetSessionID(%q) old=%q", newID, a.id) - oldID := a.id - if oldID == newID { - return nil - } - // Rename active persisted files on disk. - 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 - // Update session ID references in the system prompt. - a.r.runtime.Preamble = strings.ReplaceAll(a.r.runtime.Preamble, oldID, newID) - // Rename tmpdir so isread markers remain valid after rename. - oldTemp := filepath.Join(ollieTmpDir(), oldID) - newTemp := filepath.Join(ollieTmpDir(), newID) - if _, err := os.Stat(oldTemp); err == nil { - os.Rename(oldTemp, newTemp) //nolint:errcheck - } - a.pushSessionEnv() - return nil -} - -// defaultContextLength is used when the backend cannot report the model's -// actual context window (e.g. CodeWhisperer). 128k tokens is a safe default -// for modern models. -const defaultContextLength = 128000 -const defaultToolResultMaxBytes = 131072 - -// autoCompactLimit returns the token threshold for auto-compaction (75%). -func (a *Session) autoCompactLimit(ctx context.Context) int { - ctxLen := a.r.runtime.Backend.ContextLength(ctx) - if ctxLen <= 0 { - ctxLen = defaultContextLength - } - return ctxLen * 3 / 4 -} - -// autoWarnLimit returns the token threshold for a context-usage warning (60%). -func (a *Session) autoWarnLimit(ctx context.Context) int { - ctxLen := a.r.runtime.Backend.ContextLength(ctx) - if ctxLen <= 0 { - ctxLen = defaultContextLength - } - return ctxLen * 3 / 5 -} - -// spawnContext assembles the agent context injected at each session refresh -// point (session start, post-clear, post-compaction). It combines the -// agent-specific prompt with any agentSpawn hook output. -func (a *Session) spawnContext(ctx context.Context) string { - result := a.r.runtime.Hooks.Run(ctx, HookAgentSpawn, map[string]string{ - "session_id": a.id, - "agent": a.r.agentName, - "cwd": a.CWD(), - "model": a.r.runtime.Backend.Model(), - }, a.log) - if result.Warning != "" { - a.emit(infoEvent(result.Warning)) - } - if sum := result.Summary(); sum != "" { - a.emit(infoEvent("agentSpawn: " + sum)) - } - var parts []string - if result.Context != "" { - parts = append(parts, result.Context) - } - return strings.Join(parts, "\n\n---\n\n") -} - -// runCompact executes a full compaction cycle: pre-hook, compact, spawn-context -// re-injection, post-hook. Returns (n compacted, error). Returns (0, nil) if -// the pre-hook blocked or there was nothing to compact. Caller manages setState. -func (a *Session) runCompact(ctx context.Context, trigger string) (int, error) { - payload := map[string]string{"session_id": a.id, "trigger": trigger, "cwd": a.CWD()} - pre := a.r.runtime.Hooks.Run(ctx, HookPreCompact, payload, a.log) - if pre.Warning != "" { - a.emit(infoEvent(pre.Warning)) - } - if sum := pre.Summary(); sum != "" { - a.emit(infoEvent("preCompact: " + sum)) - } - if pre.Blocked { - a.emit(infoEvent("compact cancelled by hook")) - return 0, nil - } - if pre.Context != "" { - a.r.history.appendUserMessage(pre.Context) - } - // Use a cheaper model for compaction if configured. - compactModel := resolveCompactionModel(a.r.runtime.CompactionModel, a.r.runtime.Backend) - origModel := a.r.runtime.Backend.Model() - if compactModel != "" && compactModel != origModel { - a.r.runtime.Backend.SetModel(compactModel) - defer a.r.runtime.Backend.SetModel(origModel) - } - n, _, err := a.r.history.compact(ctx, a.r.runtime.Backend) - if err != nil { - return 0, err - } - if n > 0 { - a.auditLog.Debug("compact: removed %d messages trigger=%s session=%s", n, trigger, a.id) - a.r.warnedContext = false - if sc := a.spawnContext(ctx); sc != "" { - a.r.history.appendUserMessage(sc) - } - } - post := a.r.runtime.Hooks.Run(ctx, HookPostCompact, payload, a.log) - if post.Warning != "" { - a.emit(infoEvent(post.Warning)) - } - if sum := post.Summary(); sum != "" { - a.emit(infoEvent("postCompact: " + sum)) - } - if post.Context != "" { - a.r.history.appendUserMessage(post.Context) - } - return n, nil -} - -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() -} - -// flushSave immediately persists the session if dirty. -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 { - return - } - if a.r.history == nil || 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.history.saveToFull(path, a.id, a.r.agentName, - a.r.runtime.Backend.Name(), a.r.runtime.Backend.Model(), a.CWD(), a.remote); err != nil { - a.log.Error("session save: %v", err) - } -} - -// SaveSession writes the current session state to the given path, including -// backend and model metadata for external restore. -func (a *Session) SaveSession(path string) error { - a.mu.RLock() - defer a.mu.RUnlock() - if a.r.history == nil { - return fmt.Errorf("no active session") - } - return a.r.history.saveToFull(path, a.id, a.r.agentName, - a.r.runtime.Backend.Name(), a.r.runtime.Backend.Model(), a.CWD(), a.remote) -} - -func (a *Session) getActionCancel() context.CancelCauseFunc { - if a := a.r.currentAction.Load(); a != nil { - return a.cancel - } - return nil -} - -// Interrupt cancels the current in-progress agent turn. -// Returns true if an action was running and was cancelled. -func (a *Session) Interrupt(cause error) bool { - a.log.Debug("Interrupt() cause=%v", cause) - if cancel := a.getActionCancel(); cancel != nil { - cancel(cause) - return true - } - return false -} - -func (a *Session) Inject(prompt string) { - // If an inject is already pending, fall back to the normal FIFO so nothing - // is lost. Use CompareAndSwap to avoid a race between the nil check and store. - if !a.pendingInject.CompareAndSwap(nil, &prompt) { - a.fifo.Push(prompt) - return - } - a.emit(Event{Role: "info", Content: "\n"}) - a.emit(Event{Role: "user", Content: prompt}) -} - -func (a *Session) injectRewrite(prompt string) { - a.pendingInject.Store(&prompt) - a.emit(Event{Role: "info", Content: "\n"}) - a.emit(Event{Role: "user", Content: prompt}) -} - -func (a *Session) Queue(prompt string) { - a.fifo.Push(prompt) - a.bus.Publish("queued", prompt) -} - -func (a *Session) drainQueue() { - if prompt, ok := a.fifo.Pop(); ok { - a.Submit(context.Background(), prompt) - } -} - -func (a *Session) Bus() *pubsub.Bus { - return a.bus -} - -func (a *Session) emit(ev Event) { - a.bus.Publish("event", ev) -} - -func (a *Session) PopQueue() (string, bool) { - return a.fifo.Pop() -} - -func (a *Session) IsRunning() bool { - v := a.r.currentAction.Load() != nil - a.log.Debug("IsRunning() = %v", v) - return v -} - -func (a *Session) CtxSz() string { - if a.r.history == nil { - a.log.Debug("CtxSz() no session") - return "no active session" - } - ctxLen := a.r.runtime.Backend.ContextLength(context.Background()) - if ctxLen <= 0 { - ctxLen = defaultContextLength - } - estimated := a.r.history.estimateTokens() - pct := estimated * 100 / ctxLen - v := fmt.Sprintf("%d / %d (%d%%)", estimated, ctxLen, pct) - a.log.Debug("CtxSz() = %q", v) - return v -} - -func (a *Session) Cost() string { - if a.r.history == nil { - return "no active session" - } - return fmt.Sprintf("costLast=$%.4f\ncostSession=$%.4f\n", - a.r.history.LastTurnCostUSD, a.r.history.SessionCostUSD) -} - -func (a *Session) Usage() string { - if a.r.history == nil { - a.log.Debug("Usage() no session") - return "no active session" - } - str := fmt.Sprintf("%d in, %d out, %d requests", - a.r.history.TotalInputTokens, a.r.history.TotalOutputTokens, - a.r.history.TotalRequests) - if a.r.history.TotalCachedInputTokens > 0 { - str += fmt.Sprintf(", %d cached", a.r.history.TotalCachedInputTokens) - } - if a.r.history.Estimated { - str += " [estimated]" - } - a.log.Debug("Usage() = %q", str) - return str -} - -func (a *Session) Context() []backend.Message { - a.mu.RLock() - var msgs []backend.Message - if a.r.history != nil { - msgs = slices.Clone(a.r.history.history()) - } - a.mu.RUnlock() - if a.r.runtime.Preamble != "" { - msgs = append([]backend.Message{{Role: "system", Content: a.r.runtime.Preamble}}, msgs...) - } - return msgs -} - -func (a *Session) SystemPrompt() string { - a.log.Debug("SystemPrompt() len=%d", len(a.r.runtime.Preamble)) - return a.r.runtime.Preamble -} - -func (a *Session) GenerationParams() backend.GenerationParams { - a.mu.RLock() - defer a.mu.RUnlock() - return a.r.runtime.GenParams -} - -func (a *Session) CompactionModel() string { - a.mu.RLock() - defer a.mu.RUnlock() - return a.r.runtime.CompactionModel -} - -func (a *Session) SetCompactionModel(model string) { - a.mu.Lock() - defer a.mu.Unlock() - a.r.runtime.CompactionModel = model -} - -func (a *Session) SetGenerationParams(params backend.GenerationParams) error { - if a.IsRunning() { - return fmt.Errorf("cannot change params while agent is running") - } - a.mu.Lock() - a.r.runtime.GenParams = params - a.mu.Unlock() - return nil -} - -func (a *Session) ListModels() string { - a.log.Debug("ListModels()") - models := a.r.runtime.Backend.Models(context.Background()) - slices.Sort(models) - return strings.Join(models, "\n") -} - -// firstSentence returns the first sentence of s (up to the first period or -// newline), trimmed. Falls back to s truncated at 80 chars if no sentence end -// is found. -func firstSentence(s string) string { - for i, r := range s { - if r == '.' || r == '\n' { - return strings.TrimSpace(s[:i+1]) - } - } - if len(s) > 80 { - return s[:77] + "..." - } - return s -} - -// Submit implements Core. It processes one line of user input: slash commands -// and shell shortcuts are dispatched immediately; any other input -// starts an agent turn that streams events to the bus. If a turn is already -// in progress the prompt is queued as an in-stream interruption instead. -// -// Continuations (post-turn hook context, unconsumed inject, FIFO drain) are -// handled via an explicit loop rather than recursion to avoid stack growth. -func (a *Session) Submit(ctx context.Context, input string) { - defer func() { - if r := recover(); r != nil { - a.log.Error("panic: %v\n%s", r, debug.Stack()) - if a := a.r.currentAction.Swap(nil); a != nil { - a.cancel(fmt.Errorf("%v", r)) - } - a.setState("idle") - a.emit(Event{Role: "error", Content: fmt.Sprintf("%v", r)}) - } - }() - a.log.Debug("Submit() input_len=%d running=%v", len(input), a.IsRunning()) - if input == "" { - return - } - - // Fast path: inject and FIFO push use atomics and are safe without - // the submit lock. Handle them before acquiring submitMu so they - // don't block behind a long-running turn or command. - if a.IsRunning() { - if a.handleCommand(ctx, input) { - return - } - a.fifo.Push(input) - return - } - - // Serialize commands and turns so that e.g. a /compact arriving via - // ctl cannot race with an executeTurn arriving via prompt. - a.submitMu.Lock() - defer a.submitMu.Unlock() - - if a.handleCommand(ctx, input) { - return - } - if a.IsRunning() { - a.fifo.Push(input) - return - } - - for input != "" && ctx.Err() == nil { - input = a.executeTurn(ctx, input) - } -} - -// executeTurn runs a single agent turn and returns the next prompt to execute, -// or "" if there is nothing more to do. -func (a *Session) executeTurn(ctx context.Context, input string) string { - a.emit(Event{Role: "user", Content: input}) - - hookResult := a.r.runtime.Hooks.Run(ctx, HookPreTurn, map[string]string{ - "session_id": a.id, - "cwd": a.CWD(), - "prompt": input, - }, a.log) - if hookResult.Blocked { - a.emit(infoEvent("hook blocked prompt")) - return "" - } - if hookResult.Warning != "" { - a.emit(infoEvent(hookResult.Warning)) - } - if sum := hookResult.Summary(); sum != "" { - a.emit(infoEvent("preTurn: " + sum)) - } - if hookResult.Context != "" { - input += "\n" + hookResult.Context - } - - // Snapshot session state before this turn modifies it. Restored on failure - // so the session is clean for the next attempt. - snapSession := a.r.history - var snapMessages []backend.Message - if a.r.history != nil { - snapMessages = cloneMessages(a.r.history.messages) - } - - if a.r.history == nil { - for _, msg := range a.startupMessages { - a.log.Debug("startup: %s", msg) - a.emit(infoEvent(msg)) - } - a.startupMessages = nil - a.r.history = newHistory(input) - if sc := a.spawnContext(ctx); sc != "" { - a.r.history.appendUserMessage(sc) - } - a.r.history.appendUserMessage(input) - } else { - a.r.history.appendUserMessage(input) - } - - actCtx, actCancel := context.WithCancelCause(ctx) - handle := &actionHandle{cancel: actCancel} - a.r.currentAction.Store(handle) - a.setState("thinking") - - a.auditLog.Debug("turn: start input=%s session=%s", auditTruncate(input), a.id) - - // Build per-turn agentConfig from the current runtime. - a.r.cfg = agentConfig{ - Backend: a.r.runtime.Backend, - preamble: a.r.runtime.Preamble, - Tools: a.r.runtime.Tools, - Exec: a.r.runtime.Exec, - ClassifyTool: a.r.runtime.ClassifyTool, - ClassifyTier: a.r.runtime.ClassifyTier, - GenerationParams: a.r.runtime.GenParams, - MaxSteps: a.r.runtime.MaxSteps, - ReadPlanStep: a.readPlanStep, - TurnError: a.turnError, - } - - var replyBuf strings.Builder - a.r.cfg.Output = func(ev Event) { - switch ev.Role { - case "assistant": - replyBuf.WriteString(ev.Content) - case "call": - a.setState("calling: " + ev.Name) - a.auditLog.Debug("call: %s %s", ev.Name, auditTruncate(string(ev.Content))) - case "tool": - a.auditLog.Debug("result: %s %s", ev.Name, auditTruncate(ev.Content)) - case "state": - a.setState(ev.Content) - case "limitretry": - a.setState("limitretry") - case "error": - a.auditLog.Debug("error: %s", ev.Content) - } - if ev.Role == "usage" && a.r.history != nil { - var in, out, est, cached, creation int - var costUSD float64 - fmt.Sscanf(ev.Content, "%d %d %d %g %d %d", &in, &out, &est, &costUSD, &cached, &creation) - a.r.history.addUsage(backend.Usage{ - InputTokens: in, - CachedInputTokens: cached, - CacheCreationTokens: creation, - OutputTokens: out, - CostUSD: costUSD, - }, est != 0) - a.notifyChange() - if maxCostStr := os.Getenv("OLLIE_MAX_SESSION_COST"); maxCostStr != "" { - if limit, ferr := strconv.ParseFloat(maxCostStr, 64); ferr == nil && limit > 0 && a.r.history.SessionCostUSD >= limit { - a.emit(infoEvent(fmt.Sprintf("spending cap $%.2f reached — stopping", limit))) - a.Interrupt(ErrInterrupted) - } - } - } - a.emit(ev) - } - a.r.cfg.PopInject = func() string { - if p := a.pendingInject.Swap(nil); p != nil { - return *p - } - return "" - } - a.r.cfg.PreTool = func(ctx context.Context, name string, args json.RawMessage) HookResult { - return a.r.runtime.Hooks.Run(ctx, HookPreTool, map[string]string{ - "session_id": a.id, - "cwd": a.CWD(), - "tool": name, - "args": string(args), - }, a.log) - } - a.r.cfg.PostTool = func(ctx context.Context, name string, args json.RawMessage, result string) HookResult { - return a.r.runtime.Hooks.Run(ctx, HookPostTool, map[string]string{ - "session_id": a.id, - "cwd": a.CWD(), - "tool": name, - "args": string(args), - "result": result, - }, a.log) - } - a.r.cfg.IncrToolCallCount = func() int64 { - return a.toolCallCount.Add(1) - } - a.r.cfg.SaveSession = func() { a.saveSession() } - a.r.cfg.ResultCache = &a.r.resultCache - a.r.cfg.AutoCompact = func(ctx context.Context) { - if ctx.Err() != nil || a.r.history == nil { - return - } - limit := a.autoCompactLimit(ctx) - if limit <= 0 || a.r.history.estimateTokens() < limit { - return - } - a.emit(Event{Role: "info", Content: "auto-compacting context...\n"}) - a.setState("compacting") - if _, err := a.runCompact(ctx, "auto"); err != nil { - panic(fmt.Sprintf("mid-turn auto-compact: %v", err)) - } - a.setState("thinking") - } - - // Warn once when context usage crosses 60%; compact at 75%. - if a.r.history != nil { - tokens := a.r.history.estimateTokens() - if compactLimit := a.autoCompactLimit(ctx); compactLimit > 0 && tokens >= compactLimit { - a.emit(Event{Role: "info", Content: "auto-compacting context...\n"}) - a.setState("compacting") - if _, err := a.runCompact(ctx, "auto"); err != nil { - panic(fmt.Sprintf("auto-compact: %v", err)) - } - a.setState("thinking") - } else if warnLimit := a.autoWarnLimit(ctx); warnLimit > 0 && tokens >= warnLimit && !a.r.warnedContext { - ctxLen := a.r.cfg.Backend.ContextLength(ctx) - if ctxLen <= 0 { - ctxLen = defaultContextLength - } - pct := tokens * 100 / ctxLen - a.emit(Event{Role: "info", Content: fmt.Sprintf("context at %d%% — will auto-compact at 75%%\n", pct)}) - a.r.warnedContext = true - } - } - - // Spending cap: reject before spending more tokens. - if maxCostStr := os.Getenv("OLLIE_MAX_SESSION_COST"); maxCostStr != "" && a.r.history != nil { - if limit, err := strconv.ParseFloat(maxCostStr, 64); err == nil && limit > 0 { - if a.r.history.SessionCostUSD >= limit { - a.emit(Event{Role: "error", Content: fmt.Sprintf("spending cap $%.2f reached (session total $%.4f)", limit, a.r.history.SessionCostUSD)}) - a.setState("idle") - actCancel(nil) - a.r.currentAction.CompareAndSwap(handle, nil) - if snapSession == nil { - a.r.history = nil - } else { - a.r.history.messages = snapMessages - } - return "" - } - } - } - - if a.r.history != nil { - a.r.history.resetTurnAccumulators() - } - - // Run the turn, retrying once after compaction on context overflow. - var ( - overflowRetried bool - err error - ) - for { - err = run(actCtx, a.r.cfg, a.r.history) - actCancel(nil) - a.r.currentAction.CompareAndSwap(handle, nil) - - if err == nil { - break - } - if errors.Is(err, context.Canceled) || errors.Is(err, ErrInterrupted) { - break - } - var ctxErr *backend.ContextOverflowError - if !overflowRetried && errors.As(err, &ctxErr) && a.r.history != nil { - overflowRetried = true - a.r.history.messages = snapMessages - a.emit(Event{Role: "info", Content: "context overflow — compacting and retrying...\n"}) - a.setState("compacting") - if _, cerr := a.runCompact(ctx, "overflow"); cerr != nil { - break - } - a.r.history.appendUserMessage(input) - a.setState("thinking") - a.r.history.resetTurnAccumulators() - replyBuf.Reset() - actCtx, actCancel = context.WithCancelCause(ctx) - handle = &actionHandle{cancel: actCancel} - a.r.currentAction.Store(handle) - continue - } - break - } - - a.mu.Lock() - a.reply = replyBuf.String() - a.mu.Unlock() - replyBuf.Reset() - a.setState("idle") - a.flushSave() - - if err != nil { - // Keep completed work — only remove cancelled tool results. - // Error results are valuable feedback for the agent. - if a.r.history != nil { - a.r.history.removeCancelledToolResults() - } - if errors.Is(err, context.Canceled) || errors.Is(err, ErrInterrupted) { - a.auditLog.Debug("turn: interrupted session=%s", a.id) - a.saveSession() - return "" - } - a.emit(Event{Role: "error", Content: err.Error()}) - // Drain one FIFO item — the turnError hook may have queued a recovery prompt. - if next, ok := a.fifo.Pop(); ok { - return next - } - return "" - } - - stopResult := a.r.runtime.Hooks.Run(ctx, HookPostTurn, map[string]string{ - "session_id": a.id, - "cwd": a.CWD(), - }, a.log) - if stopResult.Warning != "" { - a.emit(infoEvent(stopResult.Warning)) - } - if sum := stopResult.Summary(); sum != "" { - a.emit(infoEvent("postTurn: " + sum)) - } - if !stopResult.Blocked && stopResult.Context != "" && a.r.history != nil { - a.r.history.appendUserMessage(stopResult.Context) - } - - if a.r.history != nil { - a.r.history.recordTurnCost(a.r.cfg.Backend.Model()) - appendUsageLog(a.id, a.r.cfg.Backend.Name(), a.r.cfg.Backend.Model(), a.r.history) - if a.r.history.LastTurnCostUSD > 0 { - a.emit(Event{Role: "info", Content: fmt.Sprintf("costLast=$%.4f\n", a.r.history.LastTurnCostUSD)}) - } - a.auditLog.Debug("turn: end reply=%s cost=$%.4f session_total=$%.4f session=%s", - auditTruncate(a.reply), a.r.history.LastTurnCostUSD, a.r.history.SessionCostUSD, a.id) - a.notifyChange() - } - a.saveSession() - - // Post-turn hook said "continue" — its context becomes the next prompt. - if stopResult.Blocked && stopResult.Context != "" { - return stopResult.Context - } - - // Inject that was pending but never consumed (text-only response with no - // tool calls) — treat it as the next user message. - if p := a.pendingInject.Swap(nil); p != nil { - return *p - } - - // Drain one item from the FIFO; the outer loop handles the rest. - if next, ok := a.fifo.Pop(); ok { - return next - } - - return "" -} - -func toolInfosToBackend(infos []tools.ToolInfo) []backend.Tool { - out := make([]backend.Tool, len(infos)) - for i, t := range infos { - out[i] = backend.Tool{ - Name: t.Name, - Description: t.Description, - Parameters: t.InputSchema, - } - } - return out -} - -func extractToolResult(raw json.RawMessage) (text string, contentBlocks []backend.ContentBlock, isError bool) { - var result struct { - IsError bool `json:"isError"` - Content []struct { - Type string `json:"type"` - Text string `json:"text"` - MediaType string `json:"media_type"` - Data string `json:"data"` - } `json:"content"` - } - if err := json.Unmarshal(raw, &result); err != nil { - return string(raw), nil, false - } - var parts []string - for _, c := range result.Content { - switch c.Type { - case "text": - parts = append(parts, c.Text) - case "image": - contentBlocks = append(contentBlocks, backend.ContentBlock{ - Type: "image", - ImageSource: &backend.ImageSource{ - Type: "base64", - MediaType: c.MediaType, - Data: c.Data, - }, - }) - } - } - return strings.Join(parts, "\n"), contentBlocks, result.IsError -} diff --git a/session/session.go b/session/session.go index bf52e6b..02d53e3 100644 --- a/session/session.go +++ b/session/session.go @@ -1,9 +1,1576 @@ package session import ( + "context" + "crypto/rand" + "encoding/json" "errors" + "fmt" + "io" + "os" + "path/filepath" + "runtime/debug" + "slices" + "strconv" + "strings" + "sync" + "sync/atomic" + "syscall" + "time" + + "github.com/simonfxr/pubsub" + "ollie/backend" + olog "ollie/log" + "ollie/paths" + "ollie/tools" ) +// toolClassifier reports whether a named tool is safe to run concurrently +// with other read-class tools. nil means treat all tools as serial. +type toolClassifier func(name string) bool + +// BuildRuntime constructs a Runtime from a pre-configured Dispatcher and +// optional agent config. cwd sets the working directory reported in the +// system prompt; if empty, the process working directory is used. +// env provides additional environment variables injected into prompt resolution +// subprocesses (e.g. OLLIE_SESSION_ID=xxx). +// The caller is responsible for registering all servers on d before calling this. +func BuildRuntime(cfg *AgentConfig, d tools.Dispatcher, cwd string, env []string, baseLayers ...string) *Runtime { + var messages []string + + var allToolInfos []tools.ToolInfo + var allTools []backend.Tool + serverOf := make(map[string]string) + + if cfg == nil || cfg.ToolsEnabled() { + var listErr error + allToolInfos, listErr = d.ListTools() + if listErr != nil { + messages = append(messages, fmt.Sprintf("list tools: %v", listErr)) + } + for _, t := range allToolInfos { + serverOf[t.Name] = t.Server + } + // Only built-in executors (with InputSchema) become backend tools. + // Named tool scripts are promoted via the tool registry and appear in the preamble. + allTools = toolInfosToBackend(allToolInfos) + + // Append named tool scripts for preamble listing only. + allToolInfos = append(allToolInfos, tools.DiscoverTools()...) + } + + hooks := Hooks{} + var preamble string + var genParams backend.GenerationParams + var maxSteps int + if cfg != nil { + for k, v := range cfg.Hooks { + hooks[k] = []string(v) + } + if resolved, err := resolvePrompt(cfg.Prompt, cwd, env); err != nil { + fmt.Fprintf(os.Stderr, "resolve prompt: %v\n", err) + } else { + preamble = resolved + } + genParams = backend.GenerationParams{ + MaxTokens: cfg.MaxTokens, + MaxCompletionTokens: cfg.MaxCompletionTokens, + Temperature: cfg.Temperature, + TopP: cfg.TopP, + TopK: cfg.TopK, + MinP: cfg.MinP, + TopA: cfg.TopA, + FrequencyPenalty: cfg.FrequencyPenalty, + PresencePenalty: cfg.PresencePenalty, + RepetitionPenalty: cfg.RepetitionPenalty, + ThinkingBudget: cfg.Reasoning, + ReasoningEffort: cfg.ReasoningEffort, + IncludeReasoning: cfg.IncludeReasoning, + ResponseFormat: cfg.ResponseFormat, + Stop: cfg.Stop, + Verbosity: cfg.Verbosity, + } + maxSteps = cfg.MaxSteps + if len(cfg.AllowTools) > 0 { + if srv, ok := d.GetServer("execute"); ok { + if rs, ok := srv.(tools.ToolRestrictionSetter); ok { + rs.SetAllowTools(cfg.AllowTools) + } + } + } + } + + exec := func(ctx context.Context, name string, args json.RawMessage) (string, []backend.ContentBlock, error) { + // Dynamically resolve the server for this tool — handles + // lazy-promoted tools that appeared after BuildRuntime. + infos, listErr := d.ListTools() + if listErr != nil { + return "", nil, listErr + } + server := "" + for _, t := range infos { + if t.Name == name { + server = t.Server + break + } + } + if server == "" { + return "", nil, fmt.Errorf("unknown tool: %s", name) + } + raw, err := d.Dispatch(ctx, server, name, args) + if err != nil { + return "", nil, err + } + text, blocks, isErr := extractToolResult(raw) + if isErr { + return "", nil, fmt.Errorf("%s", text) + } + return text, blocks, nil + } + + var classify toolClassifier + if srv, ok := d.GetServer("execute"); ok { + if pc, ok := srv.(tools.ParallelClassifier); ok { + classify = pc.IsParallelRead + } + } + + var tierFn func(string, json.RawMessage) ResultTier + if srv, ok := d.GetServer("execute"); ok { + if tc, ok := srv.(tools.TierClassifier); ok { + tierFn = func(name string, args json.RawMessage) ResultTier { + switch tc.ResultTierArgs(name, args) { + case "cold": + return TierCold + case "warm": + return TierWarm + default: + return TierHot + } + } + } + } + + var backendName, modelName, compactionModel string + if cfg != nil { + backendName = cfg.Backend + modelName = cfg.Model + compactionModel = cfg.CompactionModel + } + + // Prepend base layers (system prompt, operational model, environment) + // before the agent-specific preamble. + if len(baseLayers) > 0 { + var prefix strings.Builder + for _, layer := range baseLayers { + if layer != "" { + prefix.WriteString(layer) + prefix.WriteByte('\n') + } + } + if prefix.Len() > 0 { + preamble = prefix.String() + preamble + } + } + + // Append compact tool surface listing (name + description). + // Full prompts are available on-demand via /tools write (net/dns pattern). + var toolListing strings.Builder + for _, ti := range allToolInfos { + if ti.Description != "" && ti.Server == "" { + // Only named tool scripts (Server==""), not built-in executors + fmt.Fprintf(&toolListing, "- **%s** — %s\n", ti.Name, ti.Description) + } + } + if toolListing.Len() > 0 { + preamble += "\n# Available Tools\n\n" + toolListing.String() + } + + return &Runtime{ + Dispatcher: d, + Tools: allTools, + Exec: exec, + ClassifyTool: classify, + ClassifyTier: tierFn, + Hooks: hooks, + Preamble: preamble, + GenParams: genParams, + MaxSteps: maxSteps, + CfgBackend: backendName, + CfgModel: modelName, + CompactionModel: compactionModel, + Messages: messages, + } +} + +// DefaultPromptsDir returns the default directory for prompt templates. +func DefaultPromptsDir() string { + return paths.CfgDir() + "/prompts" +} + +// PromptsDirs returns all prompt directories from OLLIE_PROMPTS_PATH (colon-separated). +func PromptsDirs() []string { + if p := os.Getenv("OLLIE_PROMPTS_PATH"); p != "" { + return strings.Split(p, ":") + } + return []string{DefaultPromptsDir()} +} + +// AgentsDirs returns all agent directories from OLLIE_AGENTS_PATH (colon-separated). +func AgentsDirs() []string { + if p := os.Getenv("OLLIE_AGENTS_PATH"); p != "" { + return strings.Split(p, ":") + } + return []string{paths.CfgDir() + "/agents"} +} + +// AgentConfigPath resolves the config file path for a named agent. +// It searches all directories in OLLIE_AGENTS_PATH, falling back to agentsDir. +func AgentConfigPath(agentsDir, name string) string { + for _, dir := range AgentsDirs() { + p := dir + "/" + name + ".json" + if _, err := os.Stat(p); err == nil { + return p + } + } + return agentsDir + "/" + name + ".json" +} + +// NewSessionID generates a unique, lexicographically sortable session identifier. +// Format: - — sortable by creation time, unique +// even if two sessions are created within the same nanosecond. +func NewSessionID() string { + b := make([]byte, 3) + rand.Read(b) //nolint:errcheck + return strconv.FormatInt(time.Now().UnixNano(), 10) + "-" + fmt.Sprintf("%06x", b) +} + +// NewResponseID generates a unique identifier for a single assistant response. +func NewResponseID() string { + b := make([]byte, 3) + rand.Read(b) //nolint:errcheck + return "resp_" + strconv.FormatInt(time.Now().UnixNano(), 10) + "_" + fmt.Sprintf("%06x", b) +} + +// NewReactionID generates a unique identifier for a user reaction. +func NewReactionID() string { + b := make([]byte, 3) + rand.Read(b) //nolint:errcheck + return "react_" + strconv.FormatInt(time.Now().UnixNano(), 10) + "_" + fmt.Sprintf("%06x", b) +} + +// infoEvent wraps a plain-text message as an info Event. +func infoEvent(text string) Event { + return Event{Role: "info", Content: text + "\n"} +} + +// actionHandle holds the cancel function for the current agent turn. +type actionHandle struct { + cancel context.CancelCauseFunc +} + +// Config is the configuration for creating an agent. +type Config struct { + Backend backend.Backend + ModelName string // if non-empty, overrides backend's default model + AgentName string + AgentsDir string + SessionsDir string + SessionID string + Uname string // immutable user principal for 9P identity + CWD string // working directory for tool execution and system prompt + History *History + Runtime *Runtime + NewDispatcher func() tools.Dispatcher + NewBackend func(string) (backend.Backend, error) // if nil, defaults to backend.NewWithName + Log *olog.Logger // if nil, logging is disabled + // MaxSteps overrides the agent JSON maxSteps when non-zero. + // 0 means use the value from the agent config (or unlimited if absent). + MaxSteps int + // ReadPlanStep, if non-nil, is called to get the next unchecked plan step. + // When provided, it is used directly instead of reading from the filesystem. + ReadPlanStep func() string + // ListHandlers provides session-specific list implementations for slash + // commands such as /skills, /tools, /agents. Keyed by command name without + // the leading slash. Commands with no entry emit nothing. + ListHandlers map[string]func() []string + // PromptEnvExtra holds PRIME_* env vars for prompt resolution. + // For remote sessions these come from HostInfo; for local, from PromptEnv(). + // Stored so /agent reloads use the correct values. + PromptEnvExtra []string + // Remote is the SSH target for remote execution (empty = local). + Remote string + // BaseLayers are the system prompt, operational model, and environment + // block prepended to the agent preamble. Stored so /agent reloads + // preserve the server-injected context. + BaseLayers []string +} + +// Session is the concrete session type. It owns all session and agent state +// but has no knowledge of how output is rendered. +type Session struct { + // Session-level state + id string + uname string + cwd string + state string // "idle", "thinking", "calling: " + reply string + bus *pubsub.Bus + fifo Fifo + envMu sync.RWMutex + env map[string]string + plan []byte + prevPrompt string + peers map[string]bool + + // Agent + r *Agent + log *olog.Logger + sessionsDir string + readPlanStep func() string + listHandlers map[string]func() []string + turnError func(ctx context.Context, errType, errMsg string) HookResult // overridable for tests + remote string // SSH target for remote execution + startupMessages []string + toolCallCount atomic.Int64 + pendingInject atomic.Pointer[string] + mu sync.RWMutex + changeMu sync.Mutex + changeCond *sync.Cond + submitMu sync.Mutex // serializes Submit calls (commands + turns) + auditLog *olog.Logger + + // Debounced session persistence + saveMu sync.Mutex + saveDirty bool + saveTimer *time.Timer + +} + +// ToolCallCount returns the total number of tool calls executed in this +// session. The counter is monotonically increasing and never resets. +// Blocked calls (pre-tool hook exit 2) are not counted. +func (a *Session) ToolCallCount() int64 { + return a.toolCallCount.Load() +} + +// SetEnv stores a session-scoped variable and propagates it to the execute server. +func (a *Session) SetEnv(key, value string) { + a.envMu.Lock() + a.env[key] = value + a.envMu.Unlock() + if a.r.runtime == nil || a.r.runtime.Dispatcher == nil { + return + } + if srv, ok := a.r.runtime.Dispatcher.GetServer("execute"); ok { + if es, ok := srv.(tools.EnvSetter); ok { + es.SetEnv(key, value) + } + } +} + +// pushSessionEnv injects OLLIE_SESSION_ID into the execute server subprocess env. +func (a *Session) pushSessionEnv() { + if a.r.runtime == nil || a.r.runtime.Dispatcher == nil || a.id == "" { + return + } + if srv, ok := a.r.runtime.Dispatcher.GetServer("execute"); ok { + if es, ok := srv.(tools.EnvSetter); ok { + es.SetEnv("OLLIE_SESSION_ID", a.id) + if a.uname != "" { + es.SetEnv("OLLIE_UNAME", a.uname) + } + } + } +} + +// pushLockDir sets the flock directory on the execute server to the session tmpdir. +func (a *Session) pushLockDir() { + if a.r.runtime == nil || a.r.runtime.Dispatcher == nil || a.id == "" { + return + } + +} + + +var sweepTmpOnce sync.Once + +// ollieTmpDir returns the base temp directory for session tmpdirs. +func ollieTmpDir() string { + if p := os.Getenv("OLLIE_TMP_PATH"); p != "" { + return p + } + return filepath.Join(os.TempDir(), "ollie") +} + +// readNextUnchecked reads a plan file and returns the first unchecked step. +func readNextUnchecked(path string) string { + data, err := os.ReadFile(path) + if err != nil || len(data) == 0 { + return "" + } + return NextUncheckedStep(data) +} + +// 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 "" +} + +// sweepStaleTmpDirs removes session tmpdirs left by a previous crash. +// All session tmpdirs are owned by a single process, so anything +// present at startup is stale. +func sweepStaleTmpDirs() { + sweepTmpOnce.Do(func() { + base := ollieTmpDir() + os.RemoveAll(base) //nolint:errcheck + os.MkdirAll(base, 0700) //nolint:errcheck + }) +} + +// New creates an 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 = &Runtime{} + } + // Store the backend on the runtime so it's the single source of truth. + rt.Backend = cfg.Backend + // A non-zero MaxSteps in Config takes precedence over the + // value loaded from the agent JSON. + if cfg.MaxSteps > 0 { + rt.MaxSteps = cfg.MaxSteps + } + var readPlanStep func() string + if cfg.SessionID != "" { + os.MkdirAll(filepath.Join(ollieTmpDir(), cfg.SessionID), 0700) //nolint:errcheck + } + if cfg.ReadPlanStep != nil { + readPlanStep = cfg.ReadPlanStep + } + + log := cfg.Log + if log == nil { + log = olog.NewWriter("core", olog.LevelError+1, io.Discard, io.Discard) + } + + a := &Session{ + id: cfg.SessionID, + uname: cfg.Uname, + cwd: paths.ExpandHome(cfg.CWD), + state: "idle", + bus: pubsub.NewBus(), + env: make(map[string]string), + r: &Agent{ + history: cfg.History, + runtime: rt, + agentName: cfg.AgentName, + agentsDir: cfg.AgentsDir, + promptEnvExtra: cfg.PromptEnvExtra, + baseLayers: cfg.BaseLayers, + newDispatcher: cfg.NewDispatcher, + newBackend: cfg.NewBackend, + }, + log: log, + auditLog: log.Sub("audit"), + sessionsDir: cfg.SessionsDir, + remote: cfg.Remote, + startupMessages: rt.Messages, + readPlanStep: readPlanStep, + listHandlers: cfg.ListHandlers, + } + a.changeCond = sync.NewCond(&a.changeMu) + a.turnError = func(_ context.Context, errType, errMsg string) HookResult { + hookCtx, cancel := context.WithTimeout(context.Background(), time.Duration(hookTimeout)*time.Second) + defer cancel() + return a.r.runtime.Hooks.Run(hookCtx, HookTurnError, map[string]string{ + "session_id": a.id, + "cwd": a.CWD(), + "model": a.r.runtime.Backend.Model(), + "error_type": errType, + "error": errMsg, + }, a.log) + } + a.pushSessionEnv() + a.pushLockDir() + return a +} + +// Close releases resources for this session, including its tmpdir. +func (a *Session) Close() { + a.log.Debug("Close() session=%q", a.id) + a.flushSave() + if a.r.runtime != nil && a.r.runtime.Dispatcher != nil { + if srv, ok := a.r.runtime.Dispatcher.GetServer("execute"); ok { + if c, ok := srv.(interface{ Close() }); ok { + a.log.Debug("Close() calling execute.Close()") + c.Close() + } + } + } + if a.id != "" { + os.RemoveAll(filepath.Join(ollieTmpDir(), a.id)) //nolint:errcheck + } +} + +// execServer returns the execute server if available, or nil. +func (a *Session) execServer() interface{} { + if a.r.runtime == nil || a.r.runtime.Dispatcher == nil { + return nil + } + srv, _ := a.r.runtime.Dispatcher.GetServer("execute") + return srv +} + +func (a *Session) Detach() bool { + if srv := a.execServer(); srv != nil { + if d, ok := srv.(interface{ Detach() bool }); ok { + return d.Detach() + } + } + return false +} + +func (a *Session) ListDetached() []DetachedInfo { + if srv := a.execServer(); srv != nil { + type listDetacher interface { + ListDetachedRaw() []any + } + if ld, ok := srv.(listDetacher); ok { + raw := ld.ListDetachedRaw() + out := make([]DetachedInfo, 0, len(raw)) + for _, r := range raw { + if m, ok := r.(map[string]any); ok { + di := DetachedInfo{} + if v, ok := m["pid"].(int); ok { + di.PID = v + } + if v, ok := m["command"].(string); ok { + di.Command = v + } + if v, ok := m["started"].(int64); ok { + di.Started = v + } + if v, ok := m["exited"].(bool); ok { + di.Exited = v + } + if v, ok := m["exit_code"].(int); ok { + di.ExitCode = v + } + out = append(out, di) + } + } + return out + } + } + return nil +} + +func (a *Session) SignalDetached(pid, signal int) error { + if srv := a.execServer(); srv != nil { + type signaler interface { + SignalDetached(int, syscall.Signal) error + } + if sg, ok := srv.(signaler); ok { + return sg.SignalDetached(pid, syscall.Signal(signal)) + } + } + return fmt.Errorf("no execute server available") +} + +func (a *Session) GetDetachedOutput(pid int) (string, error) { + if srv := a.execServer(); srv != nil { + type outputGetter interface { + GetDetachedOutput(int) (string, error) + } + if og, ok := srv.(outputGetter); ok { + return og.GetDetachedOutput(pid) + } + } + return "", fmt.Errorf("no execute server available") +} + +func (a *Session) DismissDetached(pid int) bool { + if srv := a.execServer(); srv != nil { + type dismisser interface { + DismissDetached(int) bool + } + if d, ok := srv.(dismisser); ok { + return d.DismissDetached(pid) + } + } + return false +} + +func (a *Session) InjectSystemEvent(content string) { + a.Queue("\n" + content + "\n") +} + +// classifyReaction returns a category and description for a reaction emoji. +func classifyReaction(emoji string) (category, description string, positive bool) { + switch emoji { + case "👍", "✅": + return "positive", "The response was good. Keep doing what you're doing.", true + case "🚀", "🎉": + return "excellent", "The response was exactly what was wanted.", true + case "👎", "❌": + return "negative", "The response was wrong or unhelpful.", false + case "💩", "🤬": + return "terrible", "The response was fundamentally wrong. Stop this approach entirely and reassess from scratch.", false + case "🤔": + return "confused", "The response was unclear or confusing.", false + default: + return "unknown", "", false + } +} + +func (a *Session) Reactions() map[string]string { + result := make(map[string]string) + if a.r.history == nil { + return result + } + for _, reaction := range a.r.history.Reactions { + result[reaction.ResponseID] = reaction.Emoji + } + return result +} + +func (a *Session) React(emoji string) { + _ = a.ReactTo("", emoji) +} + +func (a *Session) ReactTo(responseID, emoji string) error { + if a.r.history == nil { + return fmt.Errorf("no active session") + } + category, _, _ := classifyReaction(emoji) + if category == "unknown" { + return fmt.Errorf("unsupported reaction: %s", emoji) + } + if responseID == "" { + for i := len(a.r.history.messages) - 1; i >= 0; i-- { + if a.r.history.messages[i].Role == "assistant" { + responseID = a.r.history.messages[i].ID + break + } + } + } + if responseID == "" { + return fmt.Errorf("no assistant response to react to") + } + found := false + for i := range a.r.history.messages { + if a.r.history.messages[i].Role == "assistant" && a.r.history.messages[i].ID == responseID { + found = true + break + } + } + if !found { + return fmt.Errorf("assistant response not found: %s", responseID) + } + + reaction := Reaction{ID: NewReactionID(), ResponseID: responseID, Emoji: emoji, Category: category, CreatedAt: time.Now()} + replaced := false + for i := range a.r.history.Reactions { + if a.r.history.Reactions[i].ResponseID == responseID { + if a.r.history.Reactions[i].Emoji == emoji { + return nil + } + a.r.history.Reactions[i] = reaction + replaced = true + break + } + } + if !replaced { + a.r.history.Reactions = append(a.r.history.Reactions, reaction) + } + a.r.history.recomputeReactionCounts() + a.saveSession() + return nil +} + +func (a *Session) AgentName() string { + v := a.r.agentName + a.log.Debug("AgentName() = %q", v) + return v +} +func (a *Session) BackendName() string { + v := a.r.runtime.Backend.Name() + a.log.Debug("BackendName() = %q", v) + return v +} +func (a *Session) ModelName() string { + v := a.r.runtime.Backend.Model() + a.log.Debug("ModelName() = %q", v) + return v +} + +func (a *Session) State() string { + return a.state +} + +func (a *Session) notifyChange() { + a.changeMu.Lock() + a.changeCond.Broadcast() + a.changeMu.Unlock() +} + +func (a *Session) setState(state string) { + a.state = state + a.log.Debug("state -> %q", state) + a.notifyChange() +} + +// WaitChange blocks until the named field changes from current, then returns +// the new value. Returns ("", false) if ctx is cancelled. +func (a *Session) WaitChange(ctx context.Context, field, current string) (string, bool) { + read := func() string { + switch field { + case WatchState: + return a.State() + case WatchUsage: + return a.Usage() + case WatchCtxSz: + return a.CtxSz() + case WatchCWD: + return a.CWD() + case WatchAgent: + return a.AgentName() + } + return "" + } + + // context.AfterFunc fires in a separate goroutine when ctx is done, + // broadcasting to unblock any waiters. + stop := context.AfterFunc(ctx, func() { + a.changeMu.Lock() + a.changeCond.Broadcast() + a.changeMu.Unlock() + }) + defer stop() + + a.changeMu.Lock() + defer a.changeMu.Unlock() + for ctx.Err() == nil { + if v := read(); v != current { + return v, true + } + a.changeCond.Wait() + } + return "", false +} + + +func (a *Session) Reply() string { + r := a.reply + a.log.Debug("Reply() len=%d", len(r)) + return r +} + +// CWD returns the current working directory for tool execution. +func (a *Session) CWD() string { + if a.cwd != "" { + a.log.Debug("CWD() = %q", a.cwd) + return a.cwd + } + wd, _ := os.Getwd() + a.log.Debug("CWD() = %q (from getwd)", wd) + return wd +} + +// SetCWD changes the working directory for tool execution and updates the +// system prompt. Returns an error if the path does not exist. +func (a *Session) SetCWD(dir string) error { + a.log.Debug("SetCWD(%q)", dir) + dir = paths.ExpandHome(dir) + if dir != "" { + if _, err := os.Stat(dir); err != nil { + return fmt.Errorf("cwd: %w", err) + } + } + oldCwd := a.cwd + a.cwd = dir + // Update cwd references in the system prompt. + if oldCwd != "" && dir != "" && oldCwd != dir { + a.r.runtime.Preamble = strings.ReplaceAll(a.r.runtime.Preamble, oldCwd, dir) + } + // Propagate to any tool server that knows how to handle it (e.g. execute). + if a.r.runtime != nil && a.r.runtime.Dispatcher != nil { + if srv, ok := a.r.runtime.Dispatcher.GetServer("execute"); ok { + if ws, ok := srv.(tools.CWDSetter); ok { + ws.SetCWD(dir) + } + } + } + a.notifyChange() + return nil +} + +// SetSessionID renames the session. It updates the in-memory ID, renames +// persisted files on disk, and propagates to the execute server env. +func (a *Session) SetSessionID(newID string) error { + a.log.Debug("SetSessionID(%q) old=%q", newID, a.id) + oldID := a.id + if oldID == newID { + return nil + } + // Rename active persisted files on disk. + 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 + // Update session ID references in the system prompt. + a.r.runtime.Preamble = strings.ReplaceAll(a.r.runtime.Preamble, oldID, newID) + // Rename tmpdir so isread markers remain valid after rename. + oldTemp := filepath.Join(ollieTmpDir(), oldID) + newTemp := filepath.Join(ollieTmpDir(), newID) + if _, err := os.Stat(oldTemp); err == nil { + os.Rename(oldTemp, newTemp) //nolint:errcheck + } + a.pushSessionEnv() + return nil +} + +// defaultContextLength is used when the backend cannot report the model's +// actual context window (e.g. CodeWhisperer). 128k tokens is a safe default +// for modern models. +const defaultContextLength = 128000 +const defaultToolResultMaxBytes = 131072 + +// autoCompactLimit returns the token threshold for auto-compaction (75%). +func (a *Session) autoCompactLimit(ctx context.Context) int { + ctxLen := a.r.runtime.Backend.ContextLength(ctx) + if ctxLen <= 0 { + ctxLen = defaultContextLength + } + return ctxLen * 3 / 4 +} + +// autoWarnLimit returns the token threshold for a context-usage warning (60%). +func (a *Session) autoWarnLimit(ctx context.Context) int { + ctxLen := a.r.runtime.Backend.ContextLength(ctx) + if ctxLen <= 0 { + ctxLen = defaultContextLength + } + return ctxLen * 3 / 5 +} + +// spawnContext assembles the agent context injected at each session refresh +// point (session start, post-clear, post-compaction). It combines the +// agent-specific prompt with any agentSpawn hook output. +func (a *Session) spawnContext(ctx context.Context) string { + result := a.r.runtime.Hooks.Run(ctx, HookAgentSpawn, map[string]string{ + "session_id": a.id, + "agent": a.r.agentName, + "cwd": a.CWD(), + "model": a.r.runtime.Backend.Model(), + }, a.log) + if result.Warning != "" { + a.emit(infoEvent(result.Warning)) + } + if sum := result.Summary(); sum != "" { + a.emit(infoEvent("agentSpawn: " + sum)) + } + var parts []string + if result.Context != "" { + parts = append(parts, result.Context) + } + return strings.Join(parts, "\n\n---\n\n") +} + +// runCompact executes a full compaction cycle: pre-hook, compact, spawn-context +// re-injection, post-hook. Returns (n compacted, error). Returns (0, nil) if +// the pre-hook blocked or there was nothing to compact. Caller manages setState. +func (a *Session) runCompact(ctx context.Context, trigger string) (int, error) { + payload := map[string]string{"session_id": a.id, "trigger": trigger, "cwd": a.CWD()} + pre := a.r.runtime.Hooks.Run(ctx, HookPreCompact, payload, a.log) + if pre.Warning != "" { + a.emit(infoEvent(pre.Warning)) + } + if sum := pre.Summary(); sum != "" { + a.emit(infoEvent("preCompact: " + sum)) + } + if pre.Blocked { + a.emit(infoEvent("compact cancelled by hook")) + return 0, nil + } + if pre.Context != "" { + a.r.history.appendUserMessage(pre.Context) + } + // Use a cheaper model for compaction if configured. + compactModel := resolveCompactionModel(a.r.runtime.CompactionModel, a.r.runtime.Backend) + origModel := a.r.runtime.Backend.Model() + if compactModel != "" && compactModel != origModel { + a.r.runtime.Backend.SetModel(compactModel) + defer a.r.runtime.Backend.SetModel(origModel) + } + n, _, err := a.r.history.compact(ctx, a.r.runtime.Backend) + if err != nil { + return 0, err + } + if n > 0 { + a.auditLog.Debug("compact: removed %d messages trigger=%s session=%s", n, trigger, a.id) + a.r.warnedContext = false + if sc := a.spawnContext(ctx); sc != "" { + a.r.history.appendUserMessage(sc) + } + } + post := a.r.runtime.Hooks.Run(ctx, HookPostCompact, payload, a.log) + if post.Warning != "" { + a.emit(infoEvent(post.Warning)) + } + if sum := post.Summary(); sum != "" { + a.emit(infoEvent("postCompact: " + sum)) + } + if post.Context != "" { + a.r.history.appendUserMessage(post.Context) + } + return n, nil +} + +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() +} + +// flushSave immediately persists the session if dirty. +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 { + return + } + if a.r.history == nil || 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.history.saveToFull(path, a.id, a.r.agentName, + a.r.runtime.Backend.Name(), a.r.runtime.Backend.Model(), a.CWD(), a.remote); err != nil { + a.log.Error("session save: %v", err) + } +} + +// SaveSession writes the current session state to the given path, including +// backend and model metadata for external restore. +func (a *Session) SaveSession(path string) error { + a.mu.RLock() + defer a.mu.RUnlock() + if a.r.history == nil { + return fmt.Errorf("no active session") + } + return a.r.history.saveToFull(path, a.id, a.r.agentName, + a.r.runtime.Backend.Name(), a.r.runtime.Backend.Model(), a.CWD(), a.remote) +} + +func (a *Session) getActionCancel() context.CancelCauseFunc { + if a := a.r.currentAction.Load(); a != nil { + return a.cancel + } + return nil +} + +// Interrupt cancels the current in-progress agent turn. +// Returns true if an action was running and was cancelled. +func (a *Session) Interrupt(cause error) bool { + a.log.Debug("Interrupt() cause=%v", cause) + if cancel := a.getActionCancel(); cancel != nil { + cancel(cause) + return true + } + return false +} + +func (a *Session) Inject(prompt string) { + // If an inject is already pending, fall back to the normal FIFO so nothing + // is lost. Use CompareAndSwap to avoid a race between the nil check and store. + if !a.pendingInject.CompareAndSwap(nil, &prompt) { + a.fifo.Push(prompt) + return + } + a.emit(Event{Role: "info", Content: "\n"}) + a.emit(Event{Role: "user", Content: prompt}) +} + +func (a *Session) injectRewrite(prompt string) { + a.pendingInject.Store(&prompt) + a.emit(Event{Role: "info", Content: "\n"}) + a.emit(Event{Role: "user", Content: prompt}) +} + +func (a *Session) Queue(prompt string) { + a.fifo.Push(prompt) + a.bus.Publish("queued", prompt) +} + +func (a *Session) drainQueue() { + if prompt, ok := a.fifo.Pop(); ok { + a.Submit(context.Background(), prompt) + } +} + +func (a *Session) Bus() *pubsub.Bus { + return a.bus +} + +func (a *Session) emit(ev Event) { + a.bus.Publish("event", ev) +} + +func (a *Session) PopQueue() (string, bool) { + return a.fifo.Pop() +} + +func (a *Session) IsRunning() bool { + v := a.r.currentAction.Load() != nil + a.log.Debug("IsRunning() = %v", v) + return v +} + +func (a *Session) CtxSz() string { + if a.r.history == nil { + a.log.Debug("CtxSz() no session") + return "no active session" + } + ctxLen := a.r.runtime.Backend.ContextLength(context.Background()) + if ctxLen <= 0 { + ctxLen = defaultContextLength + } + estimated := a.r.history.estimateTokens() + pct := estimated * 100 / ctxLen + v := fmt.Sprintf("%d / %d (%d%%)", estimated, ctxLen, pct) + a.log.Debug("CtxSz() = %q", v) + return v +} + +func (a *Session) Cost() string { + if a.r.history == nil { + return "no active session" + } + return fmt.Sprintf("costLast=$%.4f\ncostSession=$%.4f\n", + a.r.history.LastTurnCostUSD, a.r.history.SessionCostUSD) +} + +func (a *Session) Usage() string { + if a.r.history == nil { + a.log.Debug("Usage() no session") + return "no active session" + } + str := fmt.Sprintf("%d in, %d out, %d requests", + a.r.history.TotalInputTokens, a.r.history.TotalOutputTokens, + a.r.history.TotalRequests) + if a.r.history.TotalCachedInputTokens > 0 { + str += fmt.Sprintf(", %d cached", a.r.history.TotalCachedInputTokens) + } + if a.r.history.Estimated { + str += " [estimated]" + } + a.log.Debug("Usage() = %q", str) + return str +} + +func (a *Session) Context() []backend.Message { + a.mu.RLock() + var msgs []backend.Message + if a.r.history != nil { + msgs = slices.Clone(a.r.history.history()) + } + a.mu.RUnlock() + if a.r.runtime.Preamble != "" { + msgs = append([]backend.Message{{Role: "system", Content: a.r.runtime.Preamble}}, msgs...) + } + return msgs +} + +func (a *Session) SystemPrompt() string { + a.log.Debug("SystemPrompt() len=%d", len(a.r.runtime.Preamble)) + return a.r.runtime.Preamble +} + +func (a *Session) GenerationParams() backend.GenerationParams { + a.mu.RLock() + defer a.mu.RUnlock() + return a.r.runtime.GenParams +} + +func (a *Session) CompactionModel() string { + a.mu.RLock() + defer a.mu.RUnlock() + return a.r.runtime.CompactionModel +} + +func (a *Session) SetCompactionModel(model string) { + a.mu.Lock() + defer a.mu.Unlock() + a.r.runtime.CompactionModel = model +} + +func (a *Session) SetGenerationParams(params backend.GenerationParams) error { + if a.IsRunning() { + return fmt.Errorf("cannot change params while agent is running") + } + a.mu.Lock() + a.r.runtime.GenParams = params + a.mu.Unlock() + return nil +} + +func (a *Session) ListModels() string { + a.log.Debug("ListModels()") + models := a.r.runtime.Backend.Models(context.Background()) + slices.Sort(models) + return strings.Join(models, "\n") +} + +// firstSentence returns the first sentence of s (up to the first period or +// newline), trimmed. Falls back to s truncated at 80 chars if no sentence end +// is found. +func firstSentence(s string) string { + for i, r := range s { + if r == '.' || r == '\n' { + return strings.TrimSpace(s[:i+1]) + } + } + if len(s) > 80 { + return s[:77] + "..." + } + return s +} + +// Submit implements Core. It processes one line of user input: slash commands +// and shell shortcuts are dispatched immediately; any other input +// starts an agent turn that streams events to the bus. If a turn is already +// in progress the prompt is queued as an in-stream interruption instead. +// +// Continuations (post-turn hook context, unconsumed inject, FIFO drain) are +// handled via an explicit loop rather than recursion to avoid stack growth. +func (a *Session) Submit(ctx context.Context, input string) { + defer func() { + if r := recover(); r != nil { + a.log.Error("panic: %v\n%s", r, debug.Stack()) + if a := a.r.currentAction.Swap(nil); a != nil { + a.cancel(fmt.Errorf("%v", r)) + } + a.setState("idle") + a.emit(Event{Role: "error", Content: fmt.Sprintf("%v", r)}) + } + }() + a.log.Debug("Submit() input_len=%d running=%v", len(input), a.IsRunning()) + if input == "" { + return + } + + // Fast path: inject and FIFO push use atomics and are safe without + // the submit lock. Handle them before acquiring submitMu so they + // don't block behind a long-running turn or command. + if a.IsRunning() { + if a.handleCommand(ctx, input) { + return + } + a.fifo.Push(input) + return + } + + // Serialize commands and turns so that e.g. a /compact arriving via + // ctl cannot race with an executeTurn arriving via prompt. + a.submitMu.Lock() + defer a.submitMu.Unlock() + + if a.handleCommand(ctx, input) { + return + } + if a.IsRunning() { + a.fifo.Push(input) + return + } + + for input != "" && ctx.Err() == nil { + input = a.executeTurn(ctx, input) + } +} + +// executeTurn runs a single agent turn and returns the next prompt to execute, +// or "" if there is nothing more to do. +func (a *Session) executeTurn(ctx context.Context, input string) string { + a.emit(Event{Role: "user", Content: input}) + + hookResult := a.r.runtime.Hooks.Run(ctx, HookPreTurn, map[string]string{ + "session_id": a.id, + "cwd": a.CWD(), + "prompt": input, + }, a.log) + if hookResult.Blocked { + a.emit(infoEvent("hook blocked prompt")) + return "" + } + if hookResult.Warning != "" { + a.emit(infoEvent(hookResult.Warning)) + } + if sum := hookResult.Summary(); sum != "" { + a.emit(infoEvent("preTurn: " + sum)) + } + if hookResult.Context != "" { + input += "\n" + hookResult.Context + } + + // Snapshot session state before this turn modifies it. Restored on failure + // so the session is clean for the next attempt. + snapSession := a.r.history + var snapMessages []backend.Message + if a.r.history != nil { + snapMessages = cloneMessages(a.r.history.messages) + } + + if a.r.history == nil { + for _, msg := range a.startupMessages { + a.log.Debug("startup: %s", msg) + a.emit(infoEvent(msg)) + } + a.startupMessages = nil + a.r.history = newHistory(input) + if sc := a.spawnContext(ctx); sc != "" { + a.r.history.appendUserMessage(sc) + } + a.r.history.appendUserMessage(input) + } else { + a.r.history.appendUserMessage(input) + } + + actCtx, actCancel := context.WithCancelCause(ctx) + handle := &actionHandle{cancel: actCancel} + a.r.currentAction.Store(handle) + a.setState("thinking") + + a.auditLog.Debug("turn: start input=%s session=%s", auditTruncate(input), a.id) + + // Build per-turn agentConfig from the current runtime. + a.r.cfg = agentConfig{ + Backend: a.r.runtime.Backend, + preamble: a.r.runtime.Preamble, + Tools: a.r.runtime.Tools, + Exec: a.r.runtime.Exec, + ClassifyTool: a.r.runtime.ClassifyTool, + ClassifyTier: a.r.runtime.ClassifyTier, + GenerationParams: a.r.runtime.GenParams, + MaxSteps: a.r.runtime.MaxSteps, + ReadPlanStep: a.readPlanStep, + TurnError: a.turnError, + } + + var replyBuf strings.Builder + a.r.cfg.Output = func(ev Event) { + switch ev.Role { + case "assistant": + replyBuf.WriteString(ev.Content) + case "call": + a.setState("calling: " + ev.Name) + a.auditLog.Debug("call: %s %s", ev.Name, auditTruncate(string(ev.Content))) + case "tool": + a.auditLog.Debug("result: %s %s", ev.Name, auditTruncate(ev.Content)) + case "state": + a.setState(ev.Content) + case "limitretry": + a.setState("limitretry") + case "error": + a.auditLog.Debug("error: %s", ev.Content) + } + if ev.Role == "usage" && a.r.history != nil { + var in, out, est, cached, creation int + var costUSD float64 + fmt.Sscanf(ev.Content, "%d %d %d %g %d %d", &in, &out, &est, &costUSD, &cached, &creation) + a.r.history.addUsage(backend.Usage{ + InputTokens: in, + CachedInputTokens: cached, + CacheCreationTokens: creation, + OutputTokens: out, + CostUSD: costUSD, + }, est != 0) + a.notifyChange() + if maxCostStr := os.Getenv("OLLIE_MAX_SESSION_COST"); maxCostStr != "" { + if limit, ferr := strconv.ParseFloat(maxCostStr, 64); ferr == nil && limit > 0 && a.r.history.SessionCostUSD >= limit { + a.emit(infoEvent(fmt.Sprintf("spending cap $%.2f reached — stopping", limit))) + a.Interrupt(ErrInterrupted) + } + } + } + a.emit(ev) + } + a.r.cfg.PopInject = func() string { + if p := a.pendingInject.Swap(nil); p != nil { + return *p + } + return "" + } + a.r.cfg.PreTool = func(ctx context.Context, name string, args json.RawMessage) HookResult { + return a.r.runtime.Hooks.Run(ctx, HookPreTool, map[string]string{ + "session_id": a.id, + "cwd": a.CWD(), + "tool": name, + "args": string(args), + }, a.log) + } + a.r.cfg.PostTool = func(ctx context.Context, name string, args json.RawMessage, result string) HookResult { + return a.r.runtime.Hooks.Run(ctx, HookPostTool, map[string]string{ + "session_id": a.id, + "cwd": a.CWD(), + "tool": name, + "args": string(args), + "result": result, + }, a.log) + } + a.r.cfg.IncrToolCallCount = func() int64 { + return a.toolCallCount.Add(1) + } + a.r.cfg.SaveSession = func() { a.saveSession() } + a.r.cfg.ResultCache = &a.r.resultCache + a.r.cfg.AutoCompact = func(ctx context.Context) { + if ctx.Err() != nil || a.r.history == nil { + return + } + limit := a.autoCompactLimit(ctx) + if limit <= 0 || a.r.history.estimateTokens() < limit { + return + } + a.emit(Event{Role: "info", Content: "auto-compacting context...\n"}) + a.setState("compacting") + if _, err := a.runCompact(ctx, "auto"); err != nil { + panic(fmt.Sprintf("mid-turn auto-compact: %v", err)) + } + a.setState("thinking") + } + + // Warn once when context usage crosses 60%; compact at 75%. + if a.r.history != nil { + tokens := a.r.history.estimateTokens() + if compactLimit := a.autoCompactLimit(ctx); compactLimit > 0 && tokens >= compactLimit { + a.emit(Event{Role: "info", Content: "auto-compacting context...\n"}) + a.setState("compacting") + if _, err := a.runCompact(ctx, "auto"); err != nil { + panic(fmt.Sprintf("auto-compact: %v", err)) + } + a.setState("thinking") + } else if warnLimit := a.autoWarnLimit(ctx); warnLimit > 0 && tokens >= warnLimit && !a.r.warnedContext { + ctxLen := a.r.cfg.Backend.ContextLength(ctx) + if ctxLen <= 0 { + ctxLen = defaultContextLength + } + pct := tokens * 100 / ctxLen + a.emit(Event{Role: "info", Content: fmt.Sprintf("context at %d%% — will auto-compact at 75%%\n", pct)}) + a.r.warnedContext = true + } + } + + // Spending cap: reject before spending more tokens. + if maxCostStr := os.Getenv("OLLIE_MAX_SESSION_COST"); maxCostStr != "" && a.r.history != nil { + if limit, err := strconv.ParseFloat(maxCostStr, 64); err == nil && limit > 0 { + if a.r.history.SessionCostUSD >= limit { + a.emit(Event{Role: "error", Content: fmt.Sprintf("spending cap $%.2f reached (session total $%.4f)", limit, a.r.history.SessionCostUSD)}) + a.setState("idle") + actCancel(nil) + a.r.currentAction.CompareAndSwap(handle, nil) + if snapSession == nil { + a.r.history = nil + } else { + a.r.history.messages = snapMessages + } + return "" + } + } + } + + if a.r.history != nil { + a.r.history.resetTurnAccumulators() + } + + // Run the turn, retrying once after compaction on context overflow. + var ( + overflowRetried bool + err error + ) + for { + err = run(actCtx, a.r.cfg, a.r.history) + actCancel(nil) + a.r.currentAction.CompareAndSwap(handle, nil) + + if err == nil { + break + } + if errors.Is(err, context.Canceled) || errors.Is(err, ErrInterrupted) { + break + } + var ctxErr *backend.ContextOverflowError + if !overflowRetried && errors.As(err, &ctxErr) && a.r.history != nil { + overflowRetried = true + a.r.history.messages = snapMessages + a.emit(Event{Role: "info", Content: "context overflow — compacting and retrying...\n"}) + a.setState("compacting") + if _, cerr := a.runCompact(ctx, "overflow"); cerr != nil { + break + } + a.r.history.appendUserMessage(input) + a.setState("thinking") + a.r.history.resetTurnAccumulators() + replyBuf.Reset() + actCtx, actCancel = context.WithCancelCause(ctx) + handle = &actionHandle{cancel: actCancel} + a.r.currentAction.Store(handle) + continue + } + break + } + + a.mu.Lock() + a.reply = replyBuf.String() + a.mu.Unlock() + replyBuf.Reset() + a.setState("idle") + a.flushSave() + + if err != nil { + // Keep completed work — only remove cancelled tool results. + // Error results are valuable feedback for the agent. + if a.r.history != nil { + a.r.history.removeCancelledToolResults() + } + if errors.Is(err, context.Canceled) || errors.Is(err, ErrInterrupted) { + a.auditLog.Debug("turn: interrupted session=%s", a.id) + a.saveSession() + return "" + } + a.emit(Event{Role: "error", Content: err.Error()}) + // Drain one FIFO item — the turnError hook may have queued a recovery prompt. + if next, ok := a.fifo.Pop(); ok { + return next + } + return "" + } + + stopResult := a.r.runtime.Hooks.Run(ctx, HookPostTurn, map[string]string{ + "session_id": a.id, + "cwd": a.CWD(), + }, a.log) + if stopResult.Warning != "" { + a.emit(infoEvent(stopResult.Warning)) + } + if sum := stopResult.Summary(); sum != "" { + a.emit(infoEvent("postTurn: " + sum)) + } + if !stopResult.Blocked && stopResult.Context != "" && a.r.history != nil { + a.r.history.appendUserMessage(stopResult.Context) + } + + if a.r.history != nil { + a.r.history.recordTurnCost(a.r.cfg.Backend.Model()) + appendUsageLog(a.id, a.r.cfg.Backend.Name(), a.r.cfg.Backend.Model(), a.r.history) + if a.r.history.LastTurnCostUSD > 0 { + a.emit(Event{Role: "info", Content: fmt.Sprintf("costLast=$%.4f\n", a.r.history.LastTurnCostUSD)}) + } + a.auditLog.Debug("turn: end reply=%s cost=$%.4f session_total=$%.4f session=%s", + auditTruncate(a.reply), a.r.history.LastTurnCostUSD, a.r.history.SessionCostUSD, a.id) + a.notifyChange() + } + a.saveSession() + + // Post-turn hook said "continue" — its context becomes the next prompt. + if stopResult.Blocked && stopResult.Context != "" { + return stopResult.Context + } + + // Inject that was pending but never consumed (text-only response with no + // tool calls) — treat it as the next user message. + if p := a.pendingInject.Swap(nil); p != nil { + return *p + } + + // Drain one item from the FIFO; the outer loop handles the rest. + if next, ok := a.fifo.Pop(); ok { + return next + } + + return "" +} + +func toolInfosToBackend(infos []tools.ToolInfo) []backend.Tool { + out := make([]backend.Tool, len(infos)) + for i, t := range infos { + out[i] = backend.Tool{ + Name: t.Name, + Description: t.Description, + Parameters: t.InputSchema, + } + } + return out +} + +func extractToolResult(raw json.RawMessage) (text string, contentBlocks []backend.ContentBlock, isError bool) { + var result struct { + IsError bool `json:"isError"` + Content []struct { + Type string `json:"type"` + Text string `json:"text"` + MediaType string `json:"media_type"` + Data string `json:"data"` + } `json:"content"` + } + if err := json.Unmarshal(raw, &result); err != nil { + return string(raw), nil, false + } + var parts []string + for _, c := range result.Content { + switch c.Type { + case "text": + parts = append(parts, c.Text) + case "image": + contentBlocks = append(contentBlocks, backend.ContentBlock{ + Type: "image", + ImageSource: &backend.ImageSource{ + Type: "base64", + MediaType: c.MediaType, + Data: c.Data, + }, + }) + } + } + return strings.Join(parts, "\n"), contentBlocks, result.IsError +} + // WatchField names supported by Session.WaitChange. const ( WatchState = "state" diff --git a/session/waitchange_test.go b/session/waitchange_test.go index 6d30157..702ba2f 100644 --- a/session/waitchange_test.go +++ b/session/waitchange_test.go @@ -12,7 +12,7 @@ import ( olog "ollie/log" ) -// newTestCore returns a minimal harness for testing. +// newTestCore returns a minimal Session for testing. func newTestCore(initialState string) *Session { a := &Session{ id: "test",