diff options
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/ipc/ipc.go | 46 | ||||
| -rw-r--r-- | internal/ipc/server.go | 25 | ||||
| -rw-r--r-- | internal/store/store.go | 96 | ||||
| -rw-r--r-- | internal/store/store_test.go | 16 |
4 files changed, 177 insertions, 6 deletions
diff --git a/internal/ipc/ipc.go b/internal/ipc/ipc.go index dc7a1df..f5d5452 100644 --- a/internal/ipc/ipc.go +++ b/internal/ipc/ipc.go @@ -77,6 +77,18 @@ type Request struct { ClientCert *clientcert.Cert `json:"client_cert,omitempty"` ClientCertID int64 `json:"client_cert_id,omitempty"` + // For "tag_entry": ID identifies the history entry (same field + // "get"/"set_flagged"/"delete_entry" use). Plugin names who's + // tagging it - informational only, not an identity or auth + // mechanism, since anything that can reach the socket can claim any + // name. Tag is the short marker itself (e.g. "jwt", "authz-bypass"). + // Data is an opaque, plugin-defined JSON blob a panel view renders + // later without needing this plugin still connected - empty is + // fine for a plugin that only needs the tag itself, no extra detail. + TagPlugin string `json:"tag_plugin,omitempty"` + Tag string `json:"tag,omitempty"` + TagData string `json:"tag_data,omitempty"` + // For "set_flagged" and "delete_entry": ID identifies the history // entry. "clear_history" needs no fields at all. Flagged bool `json:"flagged,omitempty"` @@ -121,6 +133,9 @@ type Response struct { // per-entry insert failure is skipped, not fatal to the batch). Imported int `json:"imported,omitempty"` + // For "tag_entry": the new tag row's assigned ID. + TagID int64 `json:"tag_id,omitempty"` + // For "intrude_result": one completed attack request. IntrudeResult *IntrudeResultMsg `json:"intrude_result,omitempty"` @@ -164,6 +179,12 @@ type EntryDetail struct { // a false *Exact. RequestTruncated bool `json:"request_truncated,omitempty"` ResponseTruncated bool `json:"response_truncated,omitempty"` + // Tags is every plugin-contributed marker on this entry - see the + // "tag_entry" request. Summary.Tags (from List/Search) is just the + // comma-joined names for a compact list-view badge; this is the + // full record, including each tag's plugin and opaque Data blob, for + // a panel view to render. + Tags []store.EntryTag `json:"tags,omitempty"` } // Client talks to a mitmuxd instance for request/response queries @@ -199,6 +220,31 @@ func (c *Client) Close() error { // SetFlagged sets the flagged marker on a history entry - a simple // "mark this, revisit later" bit, filterable via flagged:true/false in // Search. +// TagEntry marks history entry id with tag, attributed to plugin (any +// non-empty name a plugin chooses to identify itself by - informational +// only), with an optional opaque data blob a panel view can render +// later. This is the core plugin-integration primitive: any process +// that can reach the control socket - the reference Go client here, or +// a plugin in any other language following the same JSON wire protocol +// (see PLAN.md) - can tag entries it finds interesting without mitmux +// needing to know anything about it in advance. Returns the new tag's +// assigned ID. +func (c *Client) TagEntry(entryID int64, plugin, tag, data string) (int64, error) { + c.mu.Lock() + defer c.mu.Unlock() + if err := c.enc.Encode(Request{Type: "tag_entry", ID: entryID, TagPlugin: plugin, Tag: tag, TagData: data}); err != nil { + return 0, err + } + var resp Response + if err := c.dec.Decode(&resp); err != nil { + return 0, err + } + if resp.Type == "error" { + return 0, errors.New(resp.Error) + } + return resp.TagID, nil +} + func (c *Client) SetFlagged(id int64, flagged bool) error { c.mu.Lock() defer c.mu.Unlock() diff --git a/internal/ipc/server.go b/internal/ipc/server.go index 378c1a9..98c2c22 100644 --- a/internal/ipc/server.go +++ b/internal/ipc/server.go @@ -141,7 +141,13 @@ func (s *Server) handleConn(conn net.Conn) { enc.Encode(Response{Type: "error", Error: err.Error()}) continue } - enc.Encode(Response{Type: "get", Detail: detailFromEntry(e)}) + detail := detailFromEntry(e) + if tags, err := s.db.ListEntryTags(req.ID); err != nil { + log.Printf("list entry tags for #%d: %v", req.ID, err) + } else { + detail.Tags = tags + } + enc.Encode(Response{Type: "get", Detail: detail}) case "ws_messages": msgs, err := s.db.ListWSMessages(req.ID) @@ -220,6 +226,23 @@ func (s *Server) handleConn(conn net.Conn) { } enc.Encode(Response{Type: "intrude_done"}) + case "tag_entry": + if req.Tag == "" { + enc.Encode(Response{Type: "error", Error: "tag_entry: missing tag"}) + continue + } + id, err := s.db.AddEntryTag(store.EntryTag{ + EntryID: req.ID, + Plugin: req.TagPlugin, + Tag: req.Tag, + Data: req.TagData, + }) + if err != nil { + enc.Encode(Response{Type: "error", Error: err.Error()}) + continue + } + enc.Encode(Response{Type: "tag_entry", TagID: id}) + case "set_flagged": if err := s.db.SetFlagged(req.ID, req.Flagged); err != nil { enc.Encode(Response{Type: "error", Error: err.Error()}) diff --git a/internal/store/store.go b/internal/store/store.go index e0f27d1..ecc3854 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -85,6 +85,24 @@ CREATE TABLE IF NOT EXISTS ws_messages ( ); CREATE INDEX IF NOT EXISTS ws_messages_entry_id ON ws_messages(entry_id); + +- Plugin-contributed markers on a history entry - see internal/ipc's +- "tag_entry" request. plugin identifies who added it (informational, +- not an identity/auth mechanism: any client connected to the socket +- can tag as anyone). data is an opaque, plugin-defined JSON blob a +- panel view can render later - e.g. a decoded JWT header/payload - +- without needing that plugin still connected. +CREATE TABLE IF NOT EXISTS entry_tags ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + entry_id INTEGER NOT NULL, + plugin TEXT NOT NULL, + tag TEXT NOT NULL, + data TEXT NOT NULL DEFAULT '', + created_at INTEGER NOT NULL +); + +CREATE INDEX IF NOT EXISTS entry_tags_entry_id ON entry_tags(entry_id); +CREATE INDEX IF NOT EXISTS entry_tags_tag ON entry_tags(tag); ` // Store is a handle to the history database. Safe for concurrent use. @@ -184,6 +202,63 @@ type Summary struct { Error string Source string Flagged bool + // Tags is every distinct plugin tag on this entry, comma-joined - + // cheap enough to compute per row (a correlated subquery, see List/ + // Search) that a plugin-tagged entry shows a badge in the history + // list itself, not just in its detail view. + Tags string +} + +// EntryTag is one plugin-contributed marker on a history entry - see +// the "tag_entry" IPC request and entry_tags' own schema comment for +// what Plugin/Data mean. +type EntryTag struct { + ID int64 + EntryID int64 + Plugin string + Tag string + Data string + CreatedAt time.Time +} + +// AddEntryTag stores t and returns its assigned ID. +func (s *Store) AddEntryTag(t EntryTag) (int64, error) { + if t.CreatedAt.IsZero() { + t.CreatedAt = time.Now() + } + res, err := s.db.Exec( + `INSERT INTO entry_tags (entry_id, plugin, tag, data, created_at) VALUES (?, ?, ?, ?, ?)`, + t.EntryID, t.Plugin, t.Tag, t.Data, t.CreatedAt.UnixMilli(), + ) + if err != nil { + return 0, fmt.Errorf("add entry tag: %w", err) + } + return res.LastInsertId() +} + +// ListEntryTags returns every tag on entryID, in the order they were +// added. +func (s *Store) ListEntryTags(entryID int64) ([]EntryTag, error) { + rows, err := s.db.Query( + `SELECT id, entry_id, plugin, tag, data, created_at FROM entry_tags WHERE entry_id = ? ORDER BY id`, + entryID, + ) + if err != nil { + return nil, fmt.Errorf("list entry tags: %w", err) + } + defer rows.Close() + + var out []EntryTag + for rows.Next() { + var t EntryTag + var createdAt int64 + if err := rows.Scan(&t.ID, &t.EntryID, &t.Plugin, &t.Tag, &t.Data, &createdAt); err != nil { + return nil, fmt.Errorf("scan entry tag row: %w", err) + } + t.CreatedAt = time.UnixMilli(createdAt) + out = append(out, t) + } + return out, rows.Err() } // Insert stores e (and indexes it for search) and returns its assigned ID. @@ -255,7 +330,8 @@ func (s *Store) List(limit int, beforeID int64) ([]Summary, error) { } rows, err := s.db.Query( `SELECT id, started_at, duration_ms, method, scheme, host, path, - COALESCE(status_code, 0), length(request_raw), COALESCE(length(response_raw), 0), error, source, flagged + COALESCE(status_code, 0), length(request_raw), COALESCE(length(response_raw), 0), error, source, flagged, + COALESCE((SELECT group_concat(DISTINCT tag) FROM entry_tags WHERE entry_tags.entry_id = history.id), '') FROM history WHERE id < ? ORDER BY id DESC LIMIT ?`, beforeID, limit, ) @@ -270,7 +346,7 @@ func (s *Store) List(limit int, beforeID int64) ([]Summary, error) { var startedAt, durationMs int64 var flagged int if err := rows.Scan(&sum.ID, &startedAt, &durationMs, &sum.Method, &sum.Scheme, &sum.Host, &sum.Path, - &sum.StatusCode, &sum.ReqSize, &sum.RespSize, &sum.Error, &sum.Source, &flagged); err != nil { + &sum.StatusCode, &sum.ReqSize, &sum.RespSize, &sum.Error, &sum.Source, &flagged, &sum.Tags); err != nil { return nil, fmt.Errorf("scan history row: %w", err) } sum.Flagged = flagged != 0 @@ -321,6 +397,12 @@ func (s *Store) Search(query string, limit int, beforeID int64) ([]Summary, erro where = append(where, "h.flagged = ?") args = append(args, boolToInt(*pred.flagged)) } + if pred.tag != "" { + where = append(where, "EXISTS (SELECT 1 FROM entry_tags WHERE entry_tags.entry_id = h.id AND entry_tags.tag = ?)") + args = append(args, pred.tag) + } + + const tagsCol = `COALESCE((SELECT group_concat(DISTINCT tag) FROM entry_tags WHERE entry_tags.entry_id = h.id), '')` var q string if remaining == "" { @@ -328,7 +410,7 @@ func (s *Store) Search(query string, limit int, beforeID int64) ([]Summary, erro // FTS5 join or ranking needed. q = `SELECT h.id, h.started_at, h.duration_ms, h.method, h.scheme, h.host, h.path, COALESCE(h.status_code, 0), length(h.request_raw), COALESCE(length(h.response_raw), 0), - h.error, h.source, h.flagged + h.error, h.source, h.flagged, ` + tagsCol + ` FROM history h WHERE ` + strings.Join(where, " AND ") + ` ORDER BY h.id DESC LIMIT ?` @@ -341,7 +423,7 @@ func (s *Store) Search(query string, limit int, beforeID int64) ([]Summary, erro args = append([]any{prepareFTSQuery(remaining)}, args...) q = `SELECT h.id, h.started_at, h.duration_ms, h.method, h.scheme, h.host, h.path, COALESCE(h.status_code, 0), length(h.request_raw), COALESCE(length(h.response_raw), 0), - h.error, h.source, h.flagged + h.error, h.source, h.flagged, ` + tagsCol + ` FROM history_fts JOIN history h ON h.id = history_fts.rowid WHERE ` + strings.Join(where, " AND ") + ` @@ -361,7 +443,7 @@ func (s *Store) Search(query string, limit int, beforeID int64) ([]Summary, erro var startedAt, durationMs int64 var flagged int if err := rows.Scan(&sum.ID, &startedAt, &durationMs, &sum.Method, &sum.Scheme, &sum.Host, &sum.Path, - &sum.StatusCode, &sum.ReqSize, &sum.RespSize, &sum.Error, &sum.Source, &flagged); err != nil { + &sum.StatusCode, &sum.ReqSize, &sum.RespSize, &sum.Error, &sum.Source, &flagged, &sum.Tags); err != nil { return nil, fmt.Errorf("scan search row: %w", err) } sum.Flagged = flagged != 0 @@ -379,6 +461,7 @@ type structuredPredicate struct { statusArgs []any source string flagged *bool + tag string } var ( @@ -427,6 +510,9 @@ func extractStructured(query string) (remaining string, pred structuredPredicate pred.flagged = &b continue } + case strings.HasPrefix(lower, "tag:"): + pred.tag = strings.TrimPrefix(f, "tag:") + continue } kept = append(kept, f) } diff --git a/internal/store/store_test.go b/internal/store/store_test.go index 7fa985a..5daa73a 100644 --- a/internal/store/store_test.go +++ b/internal/store/store_test.go @@ -13,6 +13,7 @@ func TestExtractStructured(t *testing.T) { statusArgs []any source string flagged *bool + tag string }{ { name: "plain text only", @@ -71,6 +72,18 @@ func TestExtractStructured(t *testing.T) { source: "repeater", }, { + name: "tag filter", + query: "tag:jwt", + remaining: "", + tag: "jwt", + }, + { + name: "tag filter preserves case", + query: "tag:JWT", + remaining: "", + tag: "JWT", + }, + { name: "combined with free text", query: "admin status:>=400 source:proxy", remaining: "admin", @@ -110,6 +123,9 @@ func TestExtractStructured(t *testing.T) { if pred.source != tt.source { t.Errorf("source = %q, want %q", pred.source, tt.source) } + if pred.tag != tt.tag { + t.Errorf("tag = %q, want %q", pred.tag, tt.tag) + } switch { case pred.flagged == nil && tt.flagged == nil: case pred.flagged == nil || tt.flagged == nil: |