From 5e0cb12a6b60cb00cfc045d83f280d25ee6c5de9 Mon Sep 17 00:00:00 2001 From: jolovicdev <184168873+jolovicdev@users.noreply.github.com> Date: Tue, 8 Sep 2026 08:30:28 +0200 Subject: [PATCH 1/5] Engine: scope workspaces, paginate reads, attribute sessions --- internal/engine/commands.go | 57 +++++-- internal/engine/commands_test.go | 104 +++++++++++- internal/engine/connect.go | 82 +++++++++- internal/engine/console.go | 22 ++- internal/engine/disconnect_test.go | 31 ++++ internal/engine/engine.go | 76 ++++++++- internal/engine/engine_test.go | 245 ++++++++++++++++++++++++++++ internal/engine/hailmary.go | 12 +- internal/engine/ingest.go | 80 +++++++-- internal/engine/integration_test.go | 232 ++++++++++++++++++++++++++ internal/engine/ranks.go | 24 ++- internal/engine/ranks_test.go | 49 +++++- internal/engine/refresh.go | 68 ++++++-- internal/engine/refresh_test.go | 43 +++++ internal/engine/report.go | 76 ++++++--- internal/engine/report_test.go | 72 ++++++++ internal/engine/upgrade_test.go | 6 +- internal/protocol/protocol.go | 11 +- ui/src/protocol/types.ts | 14 ++ 19 files changed, 1210 insertions(+), 94 deletions(-) diff --git a/internal/engine/commands.go b/internal/engine/commands.go index ee24ab5..4955577 100644 --- a/internal/engine/commands.go +++ b/internal/engine/commands.go @@ -15,7 +15,9 @@ func (e *Engine) execCommand(ctx context.Context, operator, method string, param if !knownMethod(method) { return nil, &protocol.ErrorBody{Code: protocol.CodeUnknownMethod, Message: "no such method: " + method} } - rpc := e.connectedRPC() + // the workspace the command starts in rides with its events: switching + // workspaces mid-RPC must not re-attribute the result + rpc, ws := e.dispatchScope() if rpc == nil { return nil, notConnected() } @@ -36,6 +38,16 @@ func (e *Engine) execCommand(ctx context.Context, operator, method string, param if err := con.Write(ctx, p.Command); err != nil { return nil, mapErr(err) } + // Readiness is restored only when a read issued after this write + // reports the console idle (consoleRead checks the generation): + // the buffered prompt must not go out here, or browsers accept + // input while msf still runs the command. + e.mu.Lock() + if e.console != nil { + e.consoleWritePending = true + e.consoleWriteGen++ + } + e.mu.Unlock() return nil, nil case protocol.MethodConsoleTabs: @@ -111,7 +123,7 @@ func (e *Engine) execCommand(ctx context.Context, operator, method string, param if eb := requireModuleRef(p.Type, p.Name); eb != nil { return nil, eb } - return e.moduleExecute(ctx, rpc, operator, p) + return e.moduleExecute(ctx, rpc, operator, ws, p) case protocol.MethodSessionAttach: var p protocol.SessionRefParams @@ -173,7 +185,7 @@ func (e *Engine) execCommand(ctx context.Context, operator, method string, param if !validPort(p.LPORT) { return nil, badParam("lport must be between 1 and 65535") } - if eb := e.sessionUpgrade(ctx, operator, p); eb != nil { + if eb := e.sessionUpgrade(ctx, operator, ws, p); eb != nil { return nil, eb } return nil, nil @@ -193,17 +205,18 @@ func (e *Engine) execCommand(ctx context.Context, operator, method string, param if p.Name == "" { return nil, badParam("name is required") } + // refreshMu spans the whole transition including the workspace RPC: + // a periodic refresh that already read the old workspace must not + // commit its stale rows after the clear below, nor interleave its + // old-workspace page reads with the new workspace's + e.refreshMu.Lock() + defer e.refreshMu.Unlock() e.mu.Lock() startGen := e.gen e.mu.Unlock() if err := gomsf.NewDbManager(rpc).SetWorkspace(ctx, p.Name); err != nil { return nil, mapErr(err) } - // refreshMu spans the whole transition: a periodic refresh that - // already read the old workspace must not commit its stale rows - // after the clear below - e.refreshMu.Lock() - defer e.refreshMu.Unlock() // msf already switched: the connection says so and the old // workspace's data must not survive a refresh that fails halfway. // The generation guard keeps a link that was replaced mid-switch @@ -254,12 +267,12 @@ func (e *Engine) execCommand(ctx context.Context, operator, method string, param if eb := parseParams(params, &p); eb != nil { return nil, eb } - return e.hailMary(ctx, operator, p) + return e.hailMary(ctx, operator, ws, p) } return nil, &protocol.ErrorBody{Code: protocol.CodeInternal, Message: "unreachable dispatch for " + method} } -func (e *Engine) moduleExecute(ctx context.Context, rpc gomsf.RPCCaller, operator string, p protocol.ModuleExecuteParams) (json.RawMessage, *protocol.ErrorBody) { +func (e *Engine) moduleExecute(ctx context.Context, rpc gomsf.RPCCaller, operator, ws string, p protocol.ModuleExecuteParams) (json.RawMessage, *protocol.ErrorBody) { if rpc == nil { return nil, notConnected() } @@ -294,15 +307,15 @@ func (e *Engine) moduleExecute(ctx context.Context, rpc gomsf.RPCCaller, operato res, err = mod.Execute(ctx) } if err != nil { - e.eventfOp(operator, protocol.LevelError, "%s/%s failed: %v", p.Type, p.Name, err) + e.eventfOpIn(ws, operator, protocol.LevelError, "%s/%s failed: %v", p.Type, p.Name, err) return nil, mapErr(err) } // msf answers job 0 when the module runs inline and finishes at once - // there is no job to watch then, and saying "job 0" just confuses if res.JobID > 0 { - e.eventfOp(operator, protocol.LevelSuccess, "%s/%s launched as job %d", p.Type, p.Name, res.JobID) + e.eventfOpIn(ws, operator, protocol.LevelSuccess, "%s/%s launched as job %d", p.Type, p.Name, res.JobID) } else { - e.eventfOp(operator, protocol.LevelSuccess, "%s/%s ran inline", p.Type, p.Name) + e.eventfOpIn(ws, operator, protocol.LevelSuccess, "%s/%s ran inline", p.Type, p.Name) } return mustJSON(protocol.ExecPayload{JobID: res.JobID, UUID: res.UUID}), nil } @@ -348,9 +361,19 @@ func (e *Engine) attach(sid string) *protocol.ErrorBody { } previous := e.interactSID e.interactSID = sid - e.interactOut = nil + if previous != sid { + // attachment is idempotent: reopening the already-attached session + // must not erase the transcript every connected operator is reading; + // that output is already consumed from the RPC stream and cannot be + // polled back + e.interactOut = nil + } mon := e.monitor - e.bus.send(protocol.InteractUpdate(&protocol.InteractState{SID: sid})) + // the update carries the buffered transcript: browsers replace their + // local copy on every interact update and would blank it without this + e.bus.send(protocol.InteractUpdate(&protocol.InteractState{ + SID: sid, Output: string(e.interactOut), + })) e.mu.Unlock() if mon != nil { if previous != "" && previous != sid { @@ -400,7 +423,7 @@ func (e *Engine) sessionWrite(ctx context.Context, p protocol.SessionWriteParams return mapErr(gomsf.NewShellSession(rpc, p.SID).Write(ctx, data)) } -func (e *Engine) sessionUpgrade(ctx context.Context, operator string, p protocol.SessionUpgradeParams) *protocol.ErrorBody { +func (e *Engine) sessionUpgrade(ctx context.Context, operator, ws string, p protocol.SessionUpgradeParams) *protocol.ErrorBody { e.mu.Lock() rpc := e.rpc session := e.sessions[p.SID] @@ -417,7 +440,7 @@ func (e *Engine) sessionUpgrade(ctx context.Context, operator string, p protocol if err := gomsf.NewShellSession(rpc, p.SID).Upgrade(ctx, p.LHOST, p.LPORT); err != nil { return mapErr(err) } - e.eventfOp(operator, protocol.LevelSuccess, "session %s upgrading to meterpreter via %s:%d", p.SID, p.LHOST, p.LPORT) + e.eventfOpIn(ws, operator, protocol.LevelSuccess, "session %s upgrading to meterpreter via %s:%d", p.SID, p.LHOST, p.LPORT) return nil } diff --git a/internal/engine/commands_test.go b/internal/engine/commands_test.go index 3829372..8235e74 100644 --- a/internal/engine/commands_test.go +++ b/internal/engine/commands_test.go @@ -5,6 +5,7 @@ import ( "encoding/json" "fmt" "sync" + "sync/atomic" "testing" "time" @@ -18,7 +19,7 @@ func TestModuleExecuteAfterDisconnectIsAnErrorNotACrash(t *testing.T) { e := connectedEngine(t, stdFake()) e.Disconnect() - _, eb := e.moduleExecute(context.Background(), e.connectedRPC(), "dana", + _, eb := e.moduleExecute(context.Background(), e.connectedRPC(), "dana", "", protocol.ModuleExecuteParams{Type: "exploit", Name: "windows/smb/a"}) if eb == nil || eb.Code != protocol.CodeNotConnected { t.Fatalf("want not_connected after disconnect, got %+v", eb) @@ -202,3 +203,104 @@ func TestWorkspaceSetSerializesWithInFlightRefresh(t *testing.T) { t.Fatalf("stale rows from the old workspace survived the switch: %d hosts", len(st.Hosts)) } } + +// Reattaching the open session keeps the transcript, and the broadcast +// carries it so browsers keep their copy too. +func TestReattachPreservesTranscript(t *testing.T) { + e := testEngine(t) + e.sessions["1"] = &protocol.SessionState{ID: "1", Type: "shell"} + if err := e.attach("1"); err != nil { + t.Fatal(err) + } + e.sessionOutput(nil, gomsf.Event{SessionID: "1", Data: "shared transcript\n"}) + sub := e.Subscribe() + defer sub.Stop() + if err := e.attach("1"); err != nil { + t.Fatal(err) + } + if got := e.State().Interact.Output; got != "shared transcript\n" { + t.Fatalf("transcript after reattach: %q", got) + } + select { + case m := <-sub.C(): + up, ok := m.(protocol.ResourceUpdate) + if !ok || up.Resource != protocol.ResInteract || up.Interact == nil { + t.Fatalf("reattach broadcast %+v", m) + } + if up.Interact.Output != "shared transcript\n" { + t.Fatalf("reattach broadcast carried %q", up.Interact.Output) + } + case <-time.After(5 * time.Second): + t.Fatal("reattach never broadcast") + } +} + +// workspace.set must not touch the remote workspace while a refresh is +// mid-fetch, and the state after the switch must be the new workspace's. +func TestWorkspaceSwitchWaitsForInFlightRefresh(t *testing.T) { + e := testEngine(t) + f := stdFake() + e.rpc = f + e.conn.Workspace = "alpha" + + var ws atomic.Value + ws.Store("alpha") + setWorkspaceEntered := make(chan struct{}) + var setOnce sync.Once + f.set(gomsf.DbSetWorkspace, func(args ...interface{}) (interface{}, error) { + ws.Store(args[0].(string)) + setOnce.Do(func() { close(setWorkspaceEntered) }) + return map[string]interface{}{"result": "success"}, nil + }) + f.set(gomsf.DbCurrentWorkspace, func(...interface{}) (interface{}, error) { + return map[string]interface{}{"workspace": ws.Load().(string)}, nil + }) + hostRead := make(chan struct{}, 1) + release := make(chan struct{}) + var once sync.Once + f.set(gomsf.DbHosts, func(args ...interface{}) (interface{}, error) { + once.Do(func() { hostRead <- struct{}{} }) + <-release + return map[string]interface{}{"hosts": []interface{}{ + map[string]interface{}{"address": args[0].(map[string]interface{})["workspace"].(string) + "-host"}, + }}, nil + }) + f.set(gomsf.DbServices, func(args ...interface{}) (interface{}, error) { + return map[string]interface{}{"services": []interface{}{ + map[string]interface{}{"host": args[0].(map[string]interface{})["workspace"].(string) + "-host", "port": 80}, + }}, nil + }) + + refreshDone := make(chan error, 1) + go func() { refreshDone <- e.refreshDB(context.Background()) }() + <-hostRead + + switchDone := make(chan struct{}) + go func() { + _, _ = e.Exec(context.Background(), "", protocol.MethodWorkspaceSet, json.RawMessage(`{"name":"beta"}`)) + close(switchDone) + }() + + select { + case <-setWorkspaceEntered: + t.Fatal("workspace.set ran during an in-flight refresh") + case <-time.After(50 * time.Millisecond): + } + + close(release) + if err := <-refreshDone; err != nil { + t.Fatal(err) + } + select { + case <-switchDone: + case <-time.After(5 * time.Second): + t.Fatal("workspace.set never finished") + } + st := e.State() + if st.Connection.Workspace != "beta" { + t.Fatalf("workspace %q", st.Connection.Workspace) + } + if len(st.Hosts) != 1 || st.Hosts[0].Address != "beta-host" { + t.Fatalf("hosts after switch: %+v", st.Hosts) + } +} diff --git a/internal/engine/connect.go b/internal/engine/connect.go index 7baab51..27e3eac 100644 --- a/internal/engine/connect.go +++ b/internal/engine/connect.go @@ -40,7 +40,14 @@ func (e *Engine) Connect(ctx context.Context, p protocol.ConnectParams) *protoco return &protocol.ErrorBody{Code: protocol.CodeBusy, Message: "already connecting or connected"} } e.connecting = true + // connectSeq names this attempt as the slot owner: a bootstrap canceled + // by Disconnect may still be unwinding when the next connect starts, and + // only the current owner may touch connecting/connectCancel + e.connectSeq++ + attempt := e.connectSeq gen := e.gen + attemptCtx, cancelAttempt := context.WithCancel(ctx) + e.connectCancel = cancelAttempt e.conn = protocol.ConnectionState{ Status: "connecting", Host: p.Host, Port: p.Port, SSL: p.SSL, Username: p.Username, @@ -48,10 +55,13 @@ func (e *Engine) Connect(ctx context.Context, p protocol.ConnectParams) *protoco e.bus.send(protocol.ConnectionUpdate(e.conn)) e.mu.Unlock() - err := e.bootstrap(ctx, p, gen) + err := e.bootstrap(attemptCtx, p, gen) e.mu.Lock() - e.connecting = false + if e.connectSeq == attempt { + e.connecting = false + e.connectCancel = nil + } ownsState := e.gen == gen if err != nil { if ownsState { // a disconnect or newer connect already owns the state @@ -60,10 +70,12 @@ func (e *Engine) Connect(ctx context.Context, p protocol.ConnectParams) *protoco e.bus.send(protocol.ConnectionUpdate(e.conn)) } e.mu.Unlock() + cancelAttempt() return &protocol.ErrorBody{Code: protocol.CodeConnectFailed, Message: err.Error()} } if e.gen != gen+1 { // bootstrap committed, then something tore the link down e.mu.Unlock() + cancelAttempt() return &protocol.ErrorBody{Code: protocol.CodeConnectFailed, Message: "connection attempt superseded"} } e.conn.Status = "connected" @@ -76,6 +88,7 @@ func (e *Engine) Connect(ctx context.Context, p protocol.ConnectParams) *protoco // that carries modules and db state to already-connected browsers e.bus.send(protocol.NewSnapshot(e.stateLocked())) e.mu.Unlock() + cancelAttempt() return nil } @@ -145,23 +158,26 @@ func (e *Engine) bootstrap(ctx context.Context, p protocol.ConnectParams, gen ui } db := gomsf.NewDbManager(rpc) - hosts, err := db.Hosts(ctx, nil) + workspace, _ := db.CurrentWorkspace(ctx) // "" when no db; not fatal + // pages are pinned to the workspace read above for the same reason the + // refresh pins its own: a console workspace switch mid-load must not mix + // collections, and msf's unpaged reads truncate at 100 rows + hosts, err := fetchAllPages(ctx, workspace, db.Hosts) if err != nil { return err } - services, err := db.Services(ctx, nil) + services, err := fetchAllPages(ctx, workspace, db.Services) if err != nil { return err } - creds, err := db.Creds(ctx, nil) + creds, err := fetchAllPages(ctx, workspace, db.Creds) if err != nil { return err } - loots, err := db.Loots(ctx, nil) + loots, err := fetchAllPages(ctx, workspace, db.Loots) if err != nil { return err } - workspace, _ := db.CurrentWorkspace(ctx) // "" when no db; not fatal e.eventf(protocol.LevelInfo, "loaded %d modules, reading database", len(modules.Exploits)+len(modules.Auxiliary)+len(modules.Post)+len(modules.Payloads)+len(modules.Encoders)+len(modules.Nops)+len(modules.Evasion)) @@ -181,6 +197,26 @@ func (e *Engine) bootstrap(ctx context.Context, p protocol.ConnectParams, gen ui } e.eventf(protocol.LevelInfo, "console ready, %d hosts and %d services in workspace", len(hosts), len(services)) + // Seed the daemon's live sessions before the monitor's first sync can + // report them as opened: sessions that already existed must not be + // re-tagged with the workspace active at (re)connect. Known + // attributions survive in sessionTags when the uuid proves the session + // is the same one - ids are reused across daemons and restarts. + // Sessions first seen now stay untagged and the report's host-membership + // fallback decides. A failed list is non-fatal: the monitor then picks + // them up tagged with the current workspace. + liveSessions, _ := gomsf.NewSessionManager(rpc).List(ctx) + e.mu.Lock() + prevTags := make(map[string]sessionTag, len(e.sessionTags)) + for sid, tag := range e.sessionTags { + prevTags[sid] = tag + } + e.mu.Unlock() + seedSessions := make(map[string]*protocol.SessionState, len(liveSessions)) + for sid, s := range liveSessions { + seedSessions[sid] = sessionState(sid, s, restoreTag(prevTags[sid], s.UUID)) + } + runCtx, cancel := context.WithCancel(context.Background()) monitor := gomsf.NewEventMonitor(runCtx, rpc, gomsf.WithEventSessionInterval(e.cfg.SessionInterval), @@ -195,7 +231,7 @@ func (e *Engine) bootstrap(ctx context.Context, p protocol.ConnectParams, gen ui e.mu.Lock() if e.gen != gen { e.mu.Unlock() - cancel() // our monitor's ctx; nobody else owns it + cancel() // our monitor's ctx; nobody else owns it return errSuperseded // the defer releases this attempt's consoles } oldCancel := e.runCancel @@ -213,6 +249,7 @@ func (e *Engine) bootstrap(ctx context.Context, p protocol.ConnectParams, gen ui e.consoleID = con.ID e.consoleOut = nil e.consolePrompt = "" + e.consoleWritePending = false e.routeConsole = routeCon e.routes = nil e.conn.MSFVersion = version.Version @@ -223,7 +260,23 @@ func (e *Engine) bootstrap(ctx context.Context, p protocol.ConnectParams, gen ui e.services = serviceStates(services) e.creds = credStates(creds) e.loot = lootStates(loots) - e.sessions = make(map[string]*protocol.SessionState) + e.sessions = seedSessions + // attribution survives the reconnect; tags of sessions that died while + // offline go with them + for sid := range e.sessionTags { + if _, live := seedSessions[sid]; !live { + delete(e.sessionTags, sid) + } + } + for sid, st := range seedSessions { + if st.Workspace == "" && st.UUID == "" { + continue + } + if e.sessionTags == nil { + e.sessionTags = make(map[string]sessionTag) + } + e.sessionTags[sid] = sessionTag{workspace: st.Workspace, uuid: st.UUID} + } e.jobs = make(map[string]*protocol.JobState) e.errStreak = 0 e.gen = gen + 1 @@ -254,6 +307,16 @@ func (e *Engine) bootstrap(ctx context.Context, p protocol.ConnectParams, gen ui func (e *Engine) Disconnect() { e.mu.Lock() e.gen++ // invalidates any bootstrap or refresh still in flight + if e.connecting { + // Free the connection slot at once: the canceled bootstrap unwinds + // on its own schedule and must not lock out a fresh connect. It + // cannot clobber the next attempt either - slot ownership is by + // connectSeq, and the old attempt fails on its canceled context. + if e.connectCancel != nil { + e.connectCancel() + } + e.connecting = false + } cancel := e.runCancel dropRPC := e.rpc dropConsoleID := e.consoleID @@ -267,6 +330,7 @@ func (e *Engine) Disconnect() { e.console = nil e.consoleID = "" e.consolePrompt = "" + e.consoleWritePending = false e.routeConsole = nil hadRoutes := len(e.routes) > 0 e.routes = nil diff --git a/internal/engine/console.go b/internal/engine/console.go index e23b207..5651576 100644 --- a/internal/engine/console.go +++ b/internal/engine/console.go @@ -17,17 +17,21 @@ func (e *Engine) consoleLoop(ctx context.Context, monitor *gomsf.EventMonitor, c case <-ctx.Done(): return case <-ticker.C: + // the generation is captured before the read is issued: a read + // that was already in flight when a write landed describes the + // old console and must not acknowledge the write's completion + gen := e.consoleGeneration() result, err := console.Read(ctx) if err != nil { e.monitorError(monitor, err) continue } - e.consoleRead(console, result) + e.consoleRead(console, result, gen) } } } -func (e *Engine) consoleRead(console *gomsf.MsfConsole, result *gomsf.ConsoleReadResult) { +func (e *Engine) consoleRead(console *gomsf.MsfConsole, result *gomsf.ConsoleReadResult, readGen uint64) { data := cleanOutput(result.Data) prompt := cleanOutput(result.Prompt) @@ -46,7 +50,11 @@ func (e *Engine) consoleRead(console *gomsf.MsfConsole, result *gomsf.ConsoleRea if data != "" { e.consoleOut = appendCapped(e.consoleOut, []byte(data)) } - if !result.Busy && prompt != "" && !bytes.HasSuffix(e.consoleOut, []byte(prompt)) { + // a stale read - issued before a still-pending write - must not append + // or emit the prompt: every prompt-terminated message reads as ready in + // browsers, whatever path it takes + staleRead := e.consoleWritePending && readGen != e.consoleWriteGen + if !result.Busy && prompt != "" && !staleRead && !bytes.HasSuffix(e.consoleOut, []byte(prompt)) { e.consoleOut = appendCapped(e.consoleOut, []byte(prompt)) stream += prompt } @@ -57,5 +65,13 @@ func (e *Engine) consoleRead(console *gomsf.MsfConsole, result *gomsf.ConsoleRea } else if stream != "" { e.bus.send(protocol.ConsoleOutputMsg{Type: protocol.KindConsoleOutput, Data: stream}) } + if !result.Busy && e.consoleWritePending && readGen == e.consoleWriteGen { + // a read issued after the write reports the console idle: browsers + // that cleared their local prompt on send get the full buffer back. + // A silent command produces no output and no prompt change, so this + // is the only path that restores them. + e.consoleWritePending = false + e.bus.send(protocol.ConsoleUpdate(&protocol.ConsoleState{Output: output})) + } e.mu.Unlock() } diff --git a/internal/engine/disconnect_test.go b/internal/engine/disconnect_test.go index 274e16e..e78d87b 100644 --- a/internal/engine/disconnect_test.go +++ b/internal/engine/disconnect_test.go @@ -1,6 +1,9 @@ package engine import ( + "context" + "errors" + "sync/atomic" "testing" "time" @@ -119,3 +122,31 @@ func TestStaleMonitorEventsDoNotResurrectClearedState(t *testing.T) { t.Fatalf("stale monitor event resurrected job: %+v", st.Jobs) } } + +// Disconnect frees the connection slot while the old bootstrap unwinds. +func TestDisconnectAllowsFreshConnect(t *testing.T) { + f := stdFake() + entered := make(chan struct{}) + release := make(chan struct{}) + var calls int32 + f.set(gomsf.CoreVersion, func(...interface{}) (interface{}, error) { + if atomic.AddInt32(&calls, 1) == 1 { + close(entered) + <-release + return nil, errors.New("old connection failed") + } + return map[string]interface{}{"version": "6.5.2"}, nil + }) + e := New(Config{RPC: f}) + t.Cleanup(e.Shutdown) + done := make(chan struct{}) + go func() { e.Connect(context.Background(), protocol.ConnectParams{Host: "old"}); close(done) }() + <-entered + e.Disconnect() + err := e.Connect(context.Background(), protocol.ConnectParams{Host: "new"}) + close(release) + <-done + if err != nil { + t.Fatalf("fresh connect after disconnect rejected: %+v", err) + } +} diff --git a/internal/engine/engine.go b/internal/engine/engine.go index 44ae014..49d4122 100644 --- a/internal/engine/engine.go +++ b/internal/engine/engine.go @@ -79,14 +79,32 @@ type Engine struct { moduleRanks map[string]string consoleOut []byte consolePrompt string - interactSID string - interactOut []byte + // consoleWritePending marks a written command whose completion no read + // has confirmed yet: readiness is restored to browsers only after a read + // reports the console idle, never straight off the write. Reads and + // writes race on the daemon, so consoleWriteGen separates reads issued + // before a write (which describe the old console) from reads issued + // after it - only the latter may acknowledge completion. + consoleWritePending bool + consoleWriteGen uint64 + interactSID string + interactOut []byte + // sessionTags remembers the workspace each live session was opened + // under, keyed by session id and validated by uuid: ids are per-daemon + // counters and get reused, so attribution carries across a reconnect + // only when the session is the same one. Survives disconnects. + sessionTags map[string]sessionTag events []*protocol.EventEntry operators map[string]int seq int64 errStreak int connecting bool + // connectSeq names the connection attempt that owns the connecting slot; + // connectCancel kills a bootstrap still unwinding after Disconnect freed + // the slot for the next attempt. Both guarded by mu. + connectSeq uint64 + connectCancel context.CancelFunc // refreshMu serializes db refreshes: the loop's periodic sweep and // command-driven ones can overlap, and a stalled read committing after // a newer refresh would revert it @@ -262,13 +280,23 @@ func (e *Engine) logf(level, format string, args ...any) { // logfOp is logf with operator attribution (team mode). Callers must hold e.mu. func (e *Engine) logfOp(operator, level, format string, args ...any) { + e.logfIn(e.conn.Workspace, operator, level, format, args...) +} + +// logfIn is logfOp with an explicit originating workspace: a command that +// started in one workspace keeps its events there even when the operator +// switches before the command's RPC returns. Callers must hold e.mu. +func (e *Engine) logfIn(ws, operator, level, format string, args ...any) { e.seq++ entry := &protocol.EventEntry{ Seq: e.seq, Time: time.Now().UTC(), Level: level, Operator: operator, - Text: fmt.Sprintf(format, args...), + // the workspace the event belongs to; empty for framework-wide + // events. Reports filter on it. + Workspace: ws, + Text: fmt.Sprintf(format, args...), } e.events = append(e.events, entry) if len(e.events) > eventRingCap { @@ -288,6 +316,14 @@ func (e *Engine) eventfOp(operator, level, format string, args ...any) { e.mu.Unlock() } +// eventfOpIn is eventfOp with an explicit originating workspace, for command +// results that land after RPC round trips. +func (e *Engine) eventfOpIn(ws, operator, level, format string, args ...any) { + e.mu.Lock() + e.logfIn(ws, operator, level, format, args...) + e.mu.Unlock() +} + // OperatorJoin and OperatorLeave track connected operators (team mode). The // operators resource fires when a name's connection count crosses zero. func (e *Engine) OperatorJoin(name string) { e.operatorDelta(name, 1) } @@ -406,6 +442,40 @@ func (e *Engine) connectedRPC() gomsf.RPCCaller { return e.rpc } +// sessionTag is the remembered attribution of one session id. +type sessionTag struct { + workspace string + uuid string +} + +// restoreTag returns the remembered workspace when the uuid matches the +// live session's. A reused id (different uuid) or a missing uuid on either +// side restores nothing: the host-membership fallback decides. +func restoreTag(tag sessionTag, uuid string) string { + if tag.uuid == "" || uuid == "" || tag.uuid != uuid { + return "" + } + return tag.workspace +} + +// dispatchScope captures the rpc handle together with the workspace a command +// starts in: the workspace rides with the command's events so a result +// landing after a workspace switch cannot be re-attributed to the new one. +func (e *Engine) dispatchScope() (gomsf.RPCCaller, string) { + e.mu.Lock() + defer e.mu.Unlock() + return e.rpc, e.conn.Workspace +} + +// consoleGeneration snapshots the write generation for a read about to be +// issued: the read may acknowledge a pending write only when the generation +// still matches, i.e. no write landed after the read was issued. +func (e *Engine) consoleGeneration() uint64 { + e.mu.Lock() + defer e.mu.Unlock() + return e.consoleWriteGen +} + var ansiRe = regexp.MustCompile(`\x1b\[[0-9;?]*[A-Za-z]`) func cleanOutput(s string) string { diff --git a/internal/engine/engine_test.go b/internal/engine/engine_test.go index e2aa3dc..cbdce08 100644 --- a/internal/engine/engine_test.go +++ b/internal/engine/engine_test.go @@ -1381,3 +1381,248 @@ func TestRefreshBroadcastsContentChanges(t *testing.T) { case <-time.After(300 * time.Millisecond): } } + +// console.write restores readiness only after a read reports the console +// idle; a silent command would otherwise leave browsers without a prompt. +func TestConsoleWriteRestoresReadinessOnlyAfterIdleRead(t *testing.T) { + e := connectedEngine(t, stdFake()) + sub := e.Subscribe() + defer sub.Stop() + + e.mu.Lock() + con := e.console + e.mu.Unlock() + e.consoleRead(con, &gomsf.ConsoleReadResult{Prompt: "msf > ", Busy: false}, e.consoleGeneration()) + // consume the seed read's consoleOutput first; bus delivery is async + select { + case m := <-sub.C(): + out, ok := m.(protocol.ConsoleOutputMsg) + if !ok || out.Data != "msf > " { + t.Fatalf("seed read broadcast %+v", m) + } + case <-time.After(5 * time.Second): + t.Fatal("seed prompt never arrived") + } + + write := func() { + if _, err := e.Exec(context.Background(), "", protocol.MethodConsoleWrite, + json.RawMessage(`{"command":"silent-command\n"}`)); err != nil { + t.Fatal(err) + } + } + expectQuiet := func() { + select { + case m := <-sub.C(): + t.Fatalf("console.write broadcast %+v before an idle post-write read", m) + default: + } + } + expectRestore := func() { + deadline := time.After(5 * time.Second) + for { + select { + case m := <-sub.C(): + up, ok := m.(protocol.ResourceUpdate) + if ok && up.Resource == protocol.ResConsole && up.Console != nil && + strings.HasSuffix(up.Console.Output, "msf > ") { + return + } + case <-deadline: + t.Fatal("no console update after the idle read") + } + } + } + + write() + expectQuiet() + e.consoleRead(con, &gomsf.ConsoleReadResult{Prompt: "msf > ", Busy: false}, e.consoleGeneration()) + expectRestore() + + // a read issued before the write describes the old console. With + // background output on it, its update must still not carry a prompt: + // browsers would read one as ready while the command runs + staleGen := e.consoleGeneration() + write() + e.consoleRead(con, &gomsf.ConsoleReadResult{Data: "background output\n", Prompt: "msf > ", Busy: false}, staleGen) + deadline := time.After(5 * time.Second) + for { + var sawUpdate bool + for { + select { + case m := <-sub.C(): + switch up := m.(type) { + case protocol.ResourceUpdate: + if up.Resource != protocol.ResConsole || up.Console == nil { + continue + } + sawUpdate = true + if strings.HasSuffix(up.Console.Output, "msf > ") { + t.Fatalf("stale read restored a ready prompt: %q", up.Console.Output) + } + case protocol.ConsoleOutputMsg: + if strings.HasSuffix(up.Data, "msf > ") { + t.Fatalf("stale read streamed a ready prompt: %q", up.Data) + } + } + case <-time.After(200 * time.Millisecond): + goto drained + } + } + drained: + if sawUpdate { + break + } + select { + case <-deadline: + t.Fatal("stale read's background output never arrived") + default: + } + } + e.consoleRead(con, &gomsf.ConsoleReadResult{Prompt: "msf > ", Busy: false}, e.consoleGeneration()) + expectRestore() +} + +// Sessions live in msfrpcd across reconnects: a reconnect under another +// workspace must not re-tag them, and sessions first seen at connect stay +// untagged so the report's host fallback decides. +func TestReconnectKeepsSessionAttribution(t *testing.T) { + f := stdFake() + var daemonSessions atomic.Value // map[string]interface{} + daemonSessions.Store(map[string]interface{}{}) + f.set(gomsf.SessionList, func(args ...interface{}) (interface{}, error) { + return daemonSessions.Load().(interface{}), nil + }) + f.set(gomsf.DbCurrentWorkspace, func(...interface{}) (interface{}, error) { + return map[string]interface{}{"workspace": "client-alpha"}, nil + }) + + e := connectedEngine(t, f) + // sessions 1 and 2 open while client-alpha is active + e.mu.Lock() + mon := e.monitor + e.mu.Unlock() + for _, sid := range []string{"1", "2"} { + e.sessionOpened(mon, gomsf.Event{SessionID: sid, + Session: &gomsf.Session{Type: "shell", UUID: "uuid-" + sid}}) + } + waitFor(t, func() bool { return len(e.State().Sessions) == 2 }) + for _, sid := range []string{"1", "2"} { + if ws := e.State().Sessions[sid].Workspace; ws != "client-alpha" { + t.Fatalf("session %s tagged %q, want client-alpha", sid, ws) + } + } + + // the link drops; session 2 dies, session 9 appears elsewhere, and the + // operator reconnects under another workspace + daemonSessions.Store(map[string]interface{}{ + "1": map[string]interface{}{"type": "shell", "target_host": "10.0.0.1", "uuid": "uuid-1"}, + "9": map[string]interface{}{"type": "meterpreter"}, + }) + f.set(gomsf.DbCurrentWorkspace, func(...interface{}) (interface{}, error) { + return map[string]interface{}{"workspace": "client-beta"}, nil + }) + e.Disconnect() + if err := e.Connect(context.Background(), protocol.ConnectParams{}); err != nil { + t.Fatalf("reconnect: %+v", err) + } + waitFor(t, func() bool { return len(e.State().Sessions) == 2 }) + sessions := e.State().Sessions + if ws := sessions["1"].Workspace; ws != "client-alpha" { + t.Fatalf("reconnect re-tagged session 1 to %q", ws) + } + if ws := sessions["9"].Workspace; ws != "" { + t.Fatalf("reconnect tagged picked-up session 9 as %q", ws) + } +} + +// Session ids are per-daemon counters: attribution carries across a +// reconnect only when the uuid proves the session is the same one. +func TestSessionAttributionValidatesUUID(t *testing.T) { + f := stdFake() + var daemonSessions atomic.Value + daemonSessions.Store(map[string]interface{}{}) + f.set(gomsf.SessionList, func(args ...interface{}) (interface{}, error) { + return daemonSessions.Load().(interface{}), nil + }) + f.set(gomsf.DbCurrentWorkspace, func(...interface{}) (interface{}, error) { + return map[string]interface{}{"workspace": "client-alpha"}, nil + }) + + e := connectedEngine(t, f) + e.mu.Lock() + mon := e.monitor + e.mu.Unlock() + e.sessionOpened(mon, gomsf.Event{SessionID: "1", + Session: &gomsf.Session{Type: "shell", UUID: "uuid-old"}}) + + // same daemon, same session: the uuid matches, attribution restores + daemonSessions.Store(map[string]interface{}{ + "1": map[string]interface{}{"type": "shell", "uuid": "uuid-old"}, + }) + e.Disconnect() + if err := e.Connect(context.Background(), protocol.ConnectParams{}); err != nil { + t.Fatalf("reconnect: %+v", err) + } + if ws := e.State().Sessions["1"].Workspace; ws != "client-alpha" { + t.Fatalf("same-uuid reconnect returned tag %q, want client-alpha", ws) + } + + // another daemon reused the id: different uuid, no attribution + daemonSessions.Store(map[string]interface{}{ + "1": map[string]interface{}{"type": "shell", "uuid": "uuid-new"}, + }) + e.Disconnect() + if err := e.Connect(context.Background(), protocol.ConnectParams{}); err != nil { + t.Fatalf("reconnect to reused id: %+v", err) + } + if ws := e.State().Sessions["1"].Workspace; ws != "" { + t.Fatalf("reused session id inherited another session's workspace %q", ws) + } +} + +// The monitor reports closes only for sessions it observed; sessions seeded +// at connect that die before the monitor's first poll are dropped by the +// reconciler. A session opened after the reconcile snapshot must survive its +// absence from that snapshot. +func TestReconcileSessionsDropsUnlistedSessions(t *testing.T) { + f := stdFake() + var daemonSessions atomic.Value + daemonSessions.Store(map[string]interface{}{ + "1": map[string]interface{}{"type": "shell", "uuid": "uuid-1"}, + }) + var e *Engine + var listCalls int32 + f.set(gomsf.SessionList, func(args ...interface{}) (interface{}, error) { + // second listing is the first reconcile: the daemon has lost + // session 1, and session 7 opens while the snapshot is taken + if atomic.AddInt32(&listCalls, 1) == 2 { + daemonSessions.Store(map[string]interface{}{ + "7": map[string]interface{}{"type": "shell", "uuid": "uuid-7"}, + }) + e.mu.Lock() + mon := e.monitor + e.mu.Unlock() + e.sessionOpened(mon, gomsf.Event{SessionID: "7", + Session: &gomsf.Session{Type: "shell", UUID: "uuid-7"}}) + return map[string]interface{}{}, nil + } + return daemonSessions.Load().(interface{}), nil + }) + e = connectedEngine(t, f) + waitFor(t, func() bool { return len(e.State().Sessions) == 1 }) + + e.reconcileSessions(context.Background()) + sessions := e.State().Sessions + if _, dead := sessions["1"]; dead { + t.Fatal("session that died before the monitor's first poll survived") + } + if _, young := sessions["7"]; !young { + t.Fatal("session opened after the reconcile snapshot was dropped") + } + + // the next sweep sees session 7 listed and keeps it + e.reconcileSessions(context.Background()) + if _, kept := e.State().Sessions["7"]; !kept { + t.Fatal("listed session dropped by reconciliation") + } +} diff --git a/internal/engine/hailmary.go b/internal/engine/hailmary.go index 6cb4228..6ea4ba6 100644 --- a/internal/engine/hailmary.go +++ b/internal/engine/hailmary.go @@ -24,7 +24,7 @@ type hailMaryTarget struct { matches []protocol.AttackMatch } -func (e *Engine) hailMary(ctx context.Context, operator string, p protocol.HailMaryParams) (json.RawMessage, *protocol.ErrorBody) { +func (e *Engine) hailMary(ctx context.Context, operator, ws string, p protocol.HailMaryParams) (json.RawMessage, *protocol.ErrorBody) { if len(p.Hosts) == 0 { return nil, &protocol.ErrorBody{Code: protocol.CodeBadParams, Message: "no hosts given"} } @@ -77,13 +77,13 @@ func (e *Engine) hailMary(ctx context.Context, operator string, p protocol.HailM launched := 0 for _, t := range targets { if len(t.matches) == 0 { - e.eventfOp(operator, protocol.LevelWarn, "hail mary: no matching exploits for %s", t.host.Address) + e.eventfOpIn(ws, operator, protocol.LevelWarn, "hail mary: no matching exploits for %s", t.host.Address) continue } - e.eventfOp(operator, protocol.LevelInfo, "hail mary on %s: launching %d exploits", t.host.Address, len(t.matches)) + e.eventfOpIn(ws, operator, protocol.LevelInfo, "hail mary on %s: launching %d exploits", t.host.Address, len(t.matches)) for _, m := range t.matches { if runCtx.Err() != nil { - e.eventfOp(operator, protocol.LevelWarn, "hail mary aborted after %d launches: connection ended", launched) + e.eventfOpIn(ws, operator, protocol.LevelWarn, "hail mary aborted after %d launches: connection ended", launched) return } options := map[string]interface{}{"RHOSTS": t.host.Address} @@ -92,7 +92,7 @@ func (e *Engine) hailMary(ctx context.Context, operator string, p protocol.HailM if m.Port > 0 { options["RPORT"] = m.Port } - if _, eb := e.moduleExecute(runCtx, e.connectedRPC(), operator, protocol.ModuleExecuteParams{ + if _, eb := e.moduleExecute(runCtx, e.connectedRPC(), operator, ws, protocol.ModuleExecuteParams{ Type: "exploit", Name: m.Name, Options: options, }); eb == nil { @@ -104,7 +104,7 @@ func (e *Engine) hailMary(ctx context.Context, operator string, p protocol.HailM } } } - e.eventfOp(operator, protocol.LevelSuccess, "hail mary finished: %d of %d planned launches", launched, planned) + e.eventfOpIn(ws, operator, protocol.LevelSuccess, "hail mary finished: %d of %d planned launches", launched, planned) }() return mustJSON(protocol.HailMaryPayload{Planned: planned}), nil diff --git a/internal/engine/ingest.go b/internal/engine/ingest.go index 963a277..c9a8199 100644 --- a/internal/engine/ingest.go +++ b/internal/engine/ingest.go @@ -1,6 +1,7 @@ package engine import ( + "context" "strings" "time" @@ -42,10 +43,31 @@ func (e *Engine) sessionOpened(m *gomsf.EventMonitor, ev gomsf.Event) { e.mu.Unlock() return } - s := ev.Session + // a session the monitor reports new belongs to the campaign that + // opened it; later workspace switches must not re-attribute it. + // Sessions seeded by bootstrap (existing at connect) never take this + // path, so they keep their original - or no - attribution. + st := sessionState(ev.SessionID, ev.Session, e.conn.Workspace) + if e.sessionTags == nil { + e.sessionTags = make(map[string]sessionTag) + } + e.sessionTags[ev.SessionID] = sessionTag{workspace: st.Workspace, uuid: st.UUID} + e.sessions[ev.SessionID] = st + sessions := copyMap(e.sessions) + host := hostLabel(st.TargetHost) + e.logf(protocol.LevelSuccess, "session %s opened (%s) on %s via %s", + ev.SessionID, st.Type, host, st.ViaExploit) + e.bus.send(protocol.SessionsUpdate(sessions)) + e.mu.Unlock() +} + +// sessionState maps a daemon session to engine state; ws is the workspace +// the session was opened under, empty when unknown. +func sessionState(id string, s *gomsf.Session, ws string) *protocol.SessionState { st := &protocol.SessionState{ - ID: ev.SessionID, - OpenedAt: time.Now().UTC(), + ID: id, + OpenedAt: time.Now().UTC(), + Workspace: ws, } if s != nil { st.Type = s.Type @@ -58,13 +80,7 @@ func (e *Engine) sessionOpened(m *gomsf.EventMonitor, ev gomsf.Event) { st.SessionHost = s.SessionHost st.UUID = s.UUID } - e.sessions[ev.SessionID] = st - sessions := copyMap(e.sessions) - host := hostLabel(st.TargetHost) - e.logf(protocol.LevelSuccess, "session %s opened (%s) on %s via %s", - ev.SessionID, st.Type, host, st.ViaExploit) - e.bus.send(protocol.SessionsUpdate(sessions)) - e.mu.Unlock() + return st } func (e *Engine) sessionClosed(m *gomsf.EventMonitor, ev gomsf.Event) { @@ -77,7 +93,7 @@ func (e *Engine) sessionClosed(m *gomsf.EventMonitor, ev gomsf.Event) { e.mu.Unlock() return } - delete(e.sessions, ev.SessionID) + e.removeSessionLocked(ev.SessionID) sessions := copyMap(e.sessions) e.logf(protocol.LevelWarn, "session %s closed", ev.SessionID) if e.interactSID == ev.SessionID { @@ -89,6 +105,48 @@ func (e *Engine) sessionClosed(m *gomsf.EventMonitor, ev gomsf.Event) { e.mu.Unlock() } +// reconcileSessions drops sessions the daemon no longer lists. The monitor +// reports closes only for sessions it observed itself, so a session seeded +// by bootstrap that dies before the monitor's first poll would stay visible +// forever. Additions stay the monitor's job: its open events carry +// attribution. +func (e *Engine) reconcileSessions(ctx context.Context) { + rpc := e.connectedRPC() + if rpc == nil { + return + } + // a session opened after this instant may be absent from the snapshot + // legitimately and is left alone this round + snapshotAt := time.Now() + live, err := gomsf.NewSessionManager(rpc).List(ctx) + if err != nil { + return + } + e.mu.Lock() + if e.rpc != rpc { + e.mu.Unlock() + return + } + changed := false + for sid, st := range e.sessions { + if _, listed := live[sid]; !listed && !st.OpenedAt.After(snapshotAt) { + e.removeSessionLocked(sid) + e.logf(protocol.LevelWarn, "session %s closed", sid) + changed = true + } + } + if changed { + e.bus.send(protocol.SessionsUpdate(copyMap(e.sessions))) + } + e.mu.Unlock() +} + +// removeSessionLocked drops a session and its attribution. Callers hold e.mu. +func (e *Engine) removeSessionLocked(sid string) { + delete(e.sessions, sid) + delete(e.sessionTags, sid) +} + func (e *Engine) sessionOutput(m *gomsf.EventMonitor, ev gomsf.Event) { e.mu.Lock() if m != e.monitor { diff --git a/internal/engine/integration_test.go b/internal/engine/integration_test.go index 7688c69..594fbb5 100644 --- a/internal/engine/integration_test.go +++ b/internal/engine/integration_test.go @@ -5,11 +5,13 @@ package engine import ( "context" "encoding/json" + "fmt" "os" "strings" "testing" "time" + "github.com/jolovicdev/go-msf/v2" "github.com/jolovicdev/hayduk/internal/protocol" ) @@ -85,3 +87,233 @@ func TestIntegrationConsoleRoundtrip(t *testing.T) { } } } + +const testWS = "hayduk-integration" + +// seedTestWorkspace fills a fresh workspace with 105 hosts and services: +// one row past the daemon's 100-row default read limit. +func seedTestWorkspace(t *testing.T) *gomsf.Client { + t.Helper() + p := integrationEnv(t) + rpc, err := gomsf.NewClient(p.Password, + gomsf.WithHost(p.Host), gomsf.WithPort(p.Port), + gomsf.WithSSL(p.SSL), gomsf.WithUsername(p.Username)) + if err != nil { + t.Fatalf("seed client login: %v", err) + } + ctx := context.Background() + _, _ = rpc.Call(ctx, gomsf.DbSetWorkspace, "default") + _, _ = rpc.Call(ctx, gomsf.DbDelWorkspace, testWS) + if _, err := rpc.Call(ctx, gomsf.DbAddWorkspace, testWS); err != nil { + t.Fatalf("add workspace: %v", err) + } + for i := 1; i <= 105; i++ { + addr := fmt.Sprintf("192.0.2.%d", i) + if _, err := rpc.Call(ctx, gomsf.DbReportHost, map[string]interface{}{ + "workspace": testWS, "host": addr, "os_name": "integration", + }); err != nil { + t.Fatalf("report_host %s: %v", addr, err) + } + if _, err := rpc.Call(ctx, gomsf.DbReportService, map[string]interface{}{ + "workspace": testWS, "host": addr, "port": 2222, "proto": "tcp", "name": "ssh", + }); err != nil { + t.Fatalf("report_service %s: %v", addr, err) + } + } + return rpc +} + +func dropTestWorkspace(rpc *gomsf.Client) { + ctx := context.Background() + _, _ = rpc.Call(ctx, gomsf.DbSetWorkspace, "default") + _, _ = rpc.Call(ctx, gomsf.DbDelWorkspace, testWS) +} + +func TestIntegrationRefreshFetchesAllPages(t *testing.T) { + seedRPC := seedTestWorkspace(t) + defer dropTestWorkspace(seedRPC) + + // the daemon caps a single unpaged read at 100 rows + raw, err := gomsf.NewDbManager(seedRPC).Hosts(context.Background(), + map[string]interface{}{"workspace": testWS}) + if err != nil { + t.Fatalf("raw hosts read: %v", err) + } + if len(raw) != 100 { + t.Fatalf("raw single read returned %d rows, want the daemon's 100 cap", len(raw)) + } + + e := New(Config{}) + t.Cleanup(e.Shutdown) + if err := e.Connect(context.Background(), integrationEnv(t)); err != nil { + t.Fatalf("connect: %+v", err) + } + if _, err := e.Exec(context.Background(), "", protocol.MethodWorkspaceSet, + json.RawMessage(`{"name":"`+testWS+`"}`)); err != nil { + t.Fatalf("workspace switch: %+v", err) + } + st := e.State() + if len(st.Hosts) != 105 || len(st.Services) != 105 { + t.Fatalf("105 rows seeded, campaign holds %d hosts and %d services", len(st.Hosts), len(st.Services)) + } + found := false + for _, h := range st.Hosts { + if h.Address == "192.0.2.105" { + found = true + } + } + if !found { + t.Fatal("host from the second page missing") + } +} + +func TestIntegrationReportWorkspaceScoped(t *testing.T) { + seedRPC := seedTestWorkspace(t) + defer dropTestWorkspace(seedRPC) + + e := New(Config{}) + t.Cleanup(e.Shutdown) + if err := e.Connect(context.Background(), integrationEnv(t)); err != nil { + t.Fatalf("connect: %+v", err) + } + + if _, err := e.Exec(context.Background(), "integration-op", protocol.MethodWorkspaceSet, + json.RawMessage(`{"name":"`+testWS+`"}`)); err != nil { + t.Fatalf("switch: %+v", err) + } + var payload protocol.ReportPayload + raw, err := e.Exec(context.Background(), "", protocol.MethodReportHTML, nil) + if err != nil { + t.Fatalf("report in %s: %+v", testWS, err) + } + if err := json.Unmarshal(raw, &payload); err != nil { + t.Fatal(err) + } + if !strings.Contains(payload.HTML, "192.0.2.105") { + t.Fatal("report missing the second-page host") + } + if !strings.Contains(payload.HTML, "integration-op") { + t.Fatal("report missing operator attribution") + } + + if _, err := e.Exec(context.Background(), "", protocol.MethodWorkspaceSet, + json.RawMessage(`{"name":"default"}`)); err != nil { + t.Fatalf("switch back: %+v", err) + } + raw, err = e.Exec(context.Background(), "", protocol.MethodReportHTML, nil) + if err != nil { + t.Fatalf("report in default: %+v", err) + } + if err := json.Unmarshal(raw, &payload); err != nil { + t.Fatal(err) + } + if strings.Contains(payload.HTML, "192.0.2.105") || strings.Contains(payload.HTML, "integration-op") { + t.Fatal("default report contains test workspace data") + } +} + +func TestIntegrationConsoleWriteAnswersConsoleState(t *testing.T) { + e := New(Config{}) + t.Cleanup(e.Shutdown) + sub := e.Subscribe() + defer sub.Stop() + if err := e.Connect(context.Background(), integrationEnv(t)); err != nil { + t.Fatalf("connect: %+v", err) + } + + if _, err := e.Exec(context.Background(), "", protocol.MethodConsoleWrite, + json.RawMessage(`{"command":"\n"}`)); err != nil { + t.Fatalf("write: %+v", err) + } + deadline := time.After(5 * time.Second) + for { + select { + case m := <-sub.C(): + up, ok := m.(protocol.ResourceUpdate) + if ok && up.Resource == protocol.ResConsole { + return + } + case <-deadline: + t.Fatal("no console update after the write") + } + } +} + +// Requires the lab's sshbox (msfadmin/msfadmin on port 2222) via +// MSF_SSH_TARGET. +func TestIntegrationSessionReattachAndWriteGuard(t *testing.T) { + target := os.Getenv("MSF_SSH_TARGET") + if target == "" { + t.Skip("MSF_SSH_TARGET not set") + } + e := New(Config{}) + t.Cleanup(e.Shutdown) + if err := e.Connect(context.Background(), integrationEnv(t)); err != nil { + t.Fatalf("connect: %+v", err) + } + + _, eb := e.Exec(context.Background(), "integration-op", protocol.MethodModuleExecute, json.RawMessage(`{ + "type":"auxiliary","name":"scanner/ssh/ssh_login", + "options":{"RHOSTS":"`+target+`","RPORT":2222,"USERNAME":"msfadmin","PASSWORD":"msfadmin","STOP_ON_SUCCESS":true} + }`)) + if eb != nil { + t.Fatalf("ssh_login launch: %+v", eb) + } + var sid string + deadline := time.After(90 * time.Second) + for sid == "" { + for s := range e.State().Sessions { + sid = s + } + if sid != "" { + break + } + select { + case <-deadline: + t.Fatal("ssh_login never opened a session") + case <-time.After(200 * time.Millisecond): + } + } + + attach := func() { + if _, eb := e.Exec(context.Background(), "", protocol.MethodSessionAttach, + json.RawMessage(`{"sid":"`+sid+`"}`)); eb != nil { + t.Fatalf("attach: %+v", eb) + } + } + attach() + if _, eb := e.Exec(context.Background(), "", protocol.MethodSessionWrite, + json.RawMessage(`{"sid":"`+sid+`","data":"echo integration-reattach\n"}`)); eb != nil { + t.Fatalf("session write: %+v", eb) + } + deadline = time.After(30 * time.Second) + for { + if strings.Contains(e.State().Interact.Output, "integration-reattach") { + break + } + select { + case <-deadline: + t.Fatalf("session output never arrived; transcript: %q", e.State().Interact.Output) + case <-time.After(100 * time.Millisecond): + } + } + + attach() + if !strings.Contains(e.State().Interact.Output, "integration-reattach") { + t.Fatalf("reattach cleared the transcript: %q", e.State().Interact.Output) + } + + if _, eb := e.Exec(context.Background(), "", protocol.MethodSessionDetach, nil); eb != nil { + t.Fatalf("detach: %+v", eb) + } + _, eb = e.Exec(context.Background(), "", protocol.MethodSessionWrite, + json.RawMessage(`{"sid":"`+sid+`","data":"echo must-not-run\n"}`)) + if eb == nil || eb.Code != protocol.CodeBusy { + t.Fatalf("write after detach: got %+v, want busy", eb) + } + + if _, eb := e.Exec(context.Background(), "", protocol.MethodSessionStop, + json.RawMessage(`{"sid":"`+sid+`"}`)); eb != nil { + t.Fatalf("stop session: %+v", eb) + } +} diff --git a/internal/engine/ranks.go b/internal/engine/ranks.go index 1fc9bf6..e9552b8 100644 --- a/internal/engine/ranks.go +++ b/internal/engine/ranks.go @@ -26,6 +26,13 @@ type rankTarget struct { modType gomsf.ModuleType } +// rankKey namespaces a rank-cache entry by module type: the exploit and +// auxiliary catalogs can contain the same relative refname, and the bare +// name would collapse both into one entry. +func rankKey(modType gomsf.ModuleType, name string) string { + return string(modType) + "/" + name +} + func (e *Engine) rankTargets() []rankTarget { e.mu.Lock() defer e.mu.Unlock() @@ -47,12 +54,13 @@ func (e *Engine) rankTargets() []rankTarget { } func (e *Engine) rankPrefetch(ctx context.Context, rpc gomsf.RPCCaller) { - // snapshot the cached names: the live map keeps mutating under flush, + // snapshot the cached keys: the live map keeps mutating under flush, // and reading it here without the lock would race cached := make(map[string]bool) e.mu.Lock() - for name := range e.moduleRanks { - cached[name] = true + gen := e.gen + for key := range e.moduleRanks { + cached[key] = true } e.mu.Unlock() @@ -70,6 +78,12 @@ func (e *Engine) rankPrefetch(ctx context.Context, rpc gomsf.RPCCaller) { batch = map[string]string{} mu.Unlock() e.mu.Lock() + if e.gen != gen { + // the link this crawl belongs to was torn down or replaced; + // committing would clobber the newer connection's cache + e.mu.Unlock() + return + } if e.moduleRanks == nil { e.moduleRanks = make(map[string]string, len(cached)+len(out)) } @@ -95,7 +109,7 @@ func (e *Engine) rankPrefetch(ctx context.Context, rpc gomsf.RPCCaller) { continue // uncached; a later connect retries } mu.Lock() - batch[t.name] = info.Rank + batch[rankKey(t.modType, t.name)] = info.Rank full := len(batch) >= rankFlushSize mu.Unlock() if full { @@ -114,7 +128,7 @@ func (e *Engine) rankPrefetch(ctx context.Context, rpc gomsf.RPCCaller) { if ctx.Err() != nil { break } - if cached[t.name] { + if cached[rankKey(t.modType, t.name)] { continue } select { diff --git a/internal/engine/ranks_test.go b/internal/engine/ranks_test.go index a849594..d84bc7a 100644 --- a/internal/engine/ranks_test.go +++ b/internal/engine/ranks_test.go @@ -35,7 +35,7 @@ func TestRankPrefetchBatchesAndCaches(t *testing.T) { waitFor(t, func() bool { return len(e.State().ModuleRanks) == 5 }) ranks := e.State().ModuleRanks - if ranks["windows/smb/a"] != "excellent" || ranks["multi/http/c"] != "normal" { + if ranks["exploit/windows/smb/a"] != "excellent" || ranks["exploit/multi/http/c"] != "normal" { t.Fatalf("ranks %+v", ranks) } @@ -59,3 +59,50 @@ func TestRankPrefetchBatchesAndCaches(t *testing.T) { } waitFor(t, func() bool { return len(e.State().ModuleRanks) == 5 }) } + +// Rank keys are namespaced by module type: exploit and auxiliary catalogs +// can contain the same refname. +func TestRankCacheSeparatesModuleTypes(t *testing.T) { + e := testEngine(t) + e.modules = &protocol.ModuleIndex{Exploits: []string{"test/shared"}, Auxiliary: []string{"test/shared"}} + f := stdFake() + f.set(gomsf.ModuleInfo, func(args ...interface{}) (interface{}, error) { + rank := "manual" + if args[0].(string) == "exploit" { + rank = "excellent" + } + return map[string]interface{}{"name": args[1], "rank": rank}, nil + }) + e.rankPrefetch(context.Background(), f) + got := e.State().ModuleRanks + if len(got) != 2 || got["exploit/test/shared"] != "excellent" || got["auxiliary/test/shared"] != "manual" { + t.Fatalf("ranks %+v, want separate per-type entries", got) + } +} + +// A stale prefetch must not commit over the replacement connection's cache. +func TestCancelledRankPrefetchCannotCommit(t *testing.T) { + e := testEngine(t) + e.modules = &protocol.ModuleIndex{Exploits: []string{"test/stale"}} + f := stdFake() + entered, resume := make(chan struct{}), make(chan struct{}) + f.set(gomsf.ModuleInfo, func(...interface{}) (interface{}, error) { + close(entered) + <-resume + return map[string]interface{}{"name": "stale", "rank": "manual"}, nil + }) + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { e.rankPrefetch(ctx, f); close(done) }() + <-entered + cancel() + e.mu.Lock() + e.gen++ + e.moduleRanks = map[string]string{"exploit/test/stale": "excellent"} + e.mu.Unlock() + close(resume) + <-done + if got := e.State().ModuleRanks["exploit/test/stale"]; got != "excellent" { + t.Fatalf("stale prefetch overwrote the rank: got %q", got) + } +} diff --git a/internal/engine/refresh.go b/internal/engine/refresh.go index c0cdc3d..0226297 100644 --- a/internal/engine/refresh.go +++ b/internal/engine/refresh.go @@ -3,6 +3,7 @@ package engine import ( "context" "errors" + "fmt" "sort" "strconv" "strings" @@ -17,12 +18,14 @@ import ( // a newer connection's context. func (e *Engine) refreshLoop(ctx context.Context) { e.refreshDB(ctx) + e.reconcileSessions(ctx) for { select { case <-ctx.Done(): return case <-time.After(e.cfg.RefreshInterval): e.refreshDB(ctx) + e.reconcileSessions(ctx) } } } @@ -44,6 +47,40 @@ func (e *Engine) refreshDB(ctx context.Context) error { return e.refreshDBLocked(ctx) } +const ( + // dbPageSize bounds one db.* RPC read. msf caps unpaged reads at its own + // default limit of 100 rows, so one naked call silently truncates every + // collection of a larger workspace. + dbPageSize = 200 + // dbMaxPages bounds the page walk so a daemon answering every page full + // (or ignoring offset entirely) cannot loop forever: 200 x 500 = 100k + // rows per collection. + dbMaxPages = 500 +) + +// fetchAllPages walks a db collection's limit/offset pages until a short page +// arrives, pinning the workspace into every request. All pages of one +// collection are fetched or none are committed: callers get an error, never a +// silently partial collection. +func fetchAllPages[T any](ctx context.Context, workspace string, fetch func(ctx context.Context, opts map[string]interface{}) ([]*T, error)) ([]*T, error) { + var out []*T + for page := 0; page < dbMaxPages; page++ { + rows, err := fetch(ctx, map[string]interface{}{ + "workspace": workspace, + "limit": dbPageSize, + "offset": page * dbPageSize, + }) + if err != nil { + return nil, err + } + out = append(out, rows...) + if len(rows) < dbPageSize { + return out, nil + } + } + return nil, fmt.Errorf("db collection exceeds %d rows", dbMaxPages*dbPageSize) +} + // refreshDBLocked is refreshDB for callers already holding refreshMu - // workspace.set uses it so its clear-plus-reload transition cannot be // interleaved with a periodic refresh that still read the old workspace. @@ -51,33 +88,40 @@ func (e *Engine) refreshDBLocked(ctx context.Context) error { e.mu.Lock() rpc := e.rpc gen := e.gen + pinned := e.conn.Workspace e.mu.Unlock() if rpc == nil { return errNotConnected } db := gomsf.NewDbManager(rpc) - hosts, err := db.Hosts(ctx, nil) + // The workspace is read once and pinned into every collection page: an + // operator switching workspaces in the raw msf console mid-refresh must + // not mix one refresh's rows across two workspaces. + if current, _ := db.CurrentWorkspace(ctx); current != "" { + pinned = current + } + hosts, err := fetchAllPages(ctx, pinned, db.Hosts) if err != nil { - e.eventf(protocol.LevelWarn, "db refresh failed: %v", err) + e.eventfOpIn(pinned, "", protocol.LevelWarn, "db refresh failed: %v", err) return err } - services, err := db.Services(ctx, nil) + services, err := fetchAllPages(ctx, pinned, db.Services) if err != nil { - e.eventf(protocol.LevelWarn, "db refresh failed: %v", err) + e.eventfOpIn(pinned, "", protocol.LevelWarn, "db refresh failed: %v", err) return err } - creds, err := db.Creds(ctx, nil) + creds, err := fetchAllPages(ctx, pinned, db.Creds) if err != nil { - e.eventf(protocol.LevelWarn, "db refresh failed: %v", err) + e.eventfOpIn(pinned, "", protocol.LevelWarn, "db refresh failed: %v", err) return err } - loots, err := db.Loots(ctx, nil) + loots, err := fetchAllPages(ctx, pinned, db.Loots) if err != nil { - e.eventf(protocol.LevelWarn, "db refresh failed: %v", err) + e.eventfOpIn(pinned, "", protocol.LevelWarn, "db refresh failed: %v", err) return err } - workspace, _ := db.CurrentWorkspace(ctx) + workspace := pinned newHosts := hostStates(hosts) newServices := serviceStates(services) @@ -119,11 +163,13 @@ func (e *Engine) refreshDBLocked(ctx context.Context) error { e.conn.Workspace = workspace e.logf(protocol.LevelInfo, "workspace changed to %s", workspace) } + // discovered rows belong to the workspace this refresh pinned, not to + // whatever the connection says by the time the commit lands for _, h := range discovered { - e.logf(protocol.LevelInfo, "discovered host %s (%s)", h.Address, hostLabel(h.Name)) + e.logfIn(workspace, "", protocol.LevelInfo, "discovered host %s (%s)", h.Address, hostLabel(h.Name)) } if credDelta > 0 { - e.logf(protocol.LevelSuccess, "%d new credentials", credDelta) + e.logfIn(workspace, "", protocol.LevelSuccess, "%d new credentials", credDelta) } if wsChanged { e.bus.send(protocol.ConnectionUpdate(e.conn)) diff --git a/internal/engine/refresh_test.go b/internal/engine/refresh_test.go index 04c82bc..c427226 100644 --- a/internal/engine/refresh_test.go +++ b/internal/engine/refresh_test.go @@ -3,7 +3,10 @@ package engine import ( "context" "errors" + "fmt" "testing" + + "github.com/jolovicdev/go-msf/v2" ) // The error streak counts consecutive failures; a successful db round trip @@ -28,3 +31,43 @@ func TestSuccessfulRefreshClearsMonitorErrorStreak(t *testing.T) { t.Fatalf("a successful db round trip must clear the error streak, got %d", streak) } } + +// msf caps unpaged db.* reads at 100 rows; the refresh walks the pages. +func TestRefreshFetchesAllDatabasePages(t *testing.T) { + e := testEngine(t) + f := stdFake() + e.rpc = f + for _, table := range []struct { + method gomsf.MsfRpcMethod + key string + }{ + {gomsf.DbHosts, "hosts"}, {gomsf.DbServices, "services"}, {gomsf.DbCreds, "creds"}, {gomsf.DbLoots, "loots"}, + } { + f.set(table.method, func(args ...interface{}) (interface{}, error) { + opts := args[0].(map[string]interface{}) + limit, offset := 100, 0 + if v, ok := opts["limit"].(int); ok { + limit = v + } + if v, ok := opts["offset"].(int); ok { + offset = v + } + rows := []interface{}{} + for i := offset; i < 105 && i < offset+limit; i++ { + rows = append(rows, map[string]interface{}{ + "address": fmt.Sprintf("10.0.0.%d", i+1), "host": "10.0.0.1", "name": fmt.Sprintf("row-%d", i), + }) + } + return map[string]interface{}{table.key: rows}, nil + }) + } + if err := e.refreshDB(context.Background()); err != nil { + t.Fatal(err) + } + s := e.State() + for name, n := range map[string]int{"hosts": len(s.Hosts), "services": len(s.Services), "creds": len(s.Creds), "loot": len(s.Loot)} { + if n != 105 { + t.Errorf("%s: 105 rows in the database, %d in the campaign", name, n) + } + } +} diff --git a/internal/engine/report.go b/internal/engine/report.go index b8ea5c2..6ea6fb5 100644 --- a/internal/engine/report.go +++ b/internal/engine/report.go @@ -132,10 +132,11 @@ var reportTmpl = template.Must(template.New("report").Parse(`
| Time | Event | ||
|---|---|---|---|
| Time | Operator | Event | |
| {{.Time}} | {{.Level}} | +{{.Operator}} | {{.Text}} |