diff options
| -rw-r--r-- | PLAN.md | 12 | ||||
| -rw-r--r-- | cmd/mitmux/main.go | 327 | ||||
| -rw-r--r-- | internal/ipc/ipc.go | 65 | ||||
| -rw-r--r-- | internal/ipc/server.go | 41 | ||||
| -rw-r--r-- | internal/proxy/proxy.go | 34 | ||||
| -rw-r--r-- | internal/rules/rules.go | 97 | ||||
| -rw-r--r-- | internal/store/store.go | 109 |
7 files changed, 681 insertions, 4 deletions
@@ -51,3 +51,15 @@ hudsucker) - same problem, worth studying even though this build is Go. boundaries, not just per-connection ones). - CA install UX per OS (Linux/macOS/Windows trust stores) - Whether WebSocket interception is v1 or a later addition +- Step 6 match-and-replace shipped headers-only. Body rules are a + separate, harder problem: request-body capture currently relies on + streaming the body straight from the client connection to the + upstream write (that's what makes it exact, byte for byte); a body + rule needs to materialize, transform, and re-send it instead, which + means deciding what "exact" even means for a rule-modified request + before touching that path again. Also still single-line-text-field + limited in the TUI (bubbles/textinput can't hold a literal CRLF), so + even with header rules, injecting a brand-new header line via the + form isn't possible yet - only rewriting/removing existing ones. The + underlying engine (rules.ApplyHeaders) already supports arbitrary + text-block edits; it's specifically the form UI that's constrained. diff --git a/cmd/mitmux/main.go b/cmd/mitmux/main.go index f47461c..68d42ed 100644 --- a/cmd/mitmux/main.go +++ b/cmd/mitmux/main.go @@ -20,6 +20,7 @@ import ( "mitmux/internal/ca" "mitmux/internal/ipc" + "mitmux/internal/rules" "mitmux/internal/store" ) @@ -71,6 +72,7 @@ const ( viewList viewMode = iota viewDetail viewRepeater + viewRules ) type detailTab int @@ -87,6 +89,16 @@ const ( focusResponse ) +type ruleField int + +const ( + fieldName ruleField = iota + fieldMatch + fieldReplace + fieldScope + fieldRegex +) + type model struct { client *ipc.Client subCh <-chan store.Summary @@ -112,6 +124,19 @@ type model struct { repeaterResult *ipc.EntryDetail sending bool + rulesTable table.Model + ruleRows []rules.Rule + + ruleForm bool // true while the add/edit form is active (vs the rule list) + ruleEditingID int64 + ruleEnabled bool + ruleName textinput.Model + ruleMatch textinput.Model + ruleReplace textinput.Model + ruleScope string // "request" or "response" + ruleRegex bool + ruleField ruleField + statusMsg string width int height int @@ -145,6 +170,24 @@ func newModel(client *ipc.Client, subCh <-chan store.Summary) *model { si.Prompt = "/" si.Placeholder = "search - plain text, or host:example.com / AND / OR / NOT" + rulesCols := []table.Column{ + {Title: "On", Width: 3}, + {Title: "Name", Width: 16}, + {Title: "Scope", Width: 9}, + {Title: "Match", Width: 24}, + {Title: "Replace", Width: 24}, + {Title: "Regex", Width: 5}, + } + rt := table.New(table.WithColumns(rulesCols), table.WithFocused(true)) + rt.SetStyles(st) + + nameIn := textinput.New() + nameIn.Placeholder = "rule name" + matchIn := textinput.New() + matchIn.Placeholder = "match text or regex" + replaceIn := textinput.New() + replaceIn.Placeholder = "replacement" + return &model{ client: client, subCh: subCh, @@ -152,6 +195,11 @@ func newModel(client *ipc.Client, subCh <-chan store.Summary) *model { table: t, reqArea: ta, searchInput: si, + rulesTable: rt, + ruleName: nameIn, + ruleMatch: matchIn, + ruleReplace: replaceIn, + ruleScope: "request", } } @@ -227,6 +275,106 @@ func (m *model) enterRepeater(d *ipc.EntryDetail) { m.statusMsg = "" } +type rulesLoadedMsg struct { + rules []rules.Rule + err error +} + +type ruleWriteDoneMsg struct { + action string // "saved", "deleted", "toggled" - for the status line + err error +} + +func (m *model) loadRules() tea.Msg { + rs, err := m.client.ListRules() + return rulesLoadedMsg{rules: rs, err: err} +} + +func (m *model) saveRule(r rules.Rule) tea.Cmd { + return func() tea.Msg { + _, err := m.client.SaveRule(r) + return ruleWriteDoneMsg{action: "saved", err: err} + } +} + +func (m *model) deleteSelectedRule() tea.Cmd { + row := m.rulesTable.Cursor() + if row < 0 || row >= len(m.ruleRows) { + return nil + } + id := m.ruleRows[row].ID + return func() tea.Msg { + err := m.client.DeleteRule(id) + return ruleWriteDoneMsg{action: "deleted", err: err} + } +} + +func (m *model) toggleSelectedRule() tea.Cmd { + row := m.rulesTable.Cursor() + if row < 0 || row >= len(m.ruleRows) { + return nil + } + r := m.ruleRows[row] + return func() tea.Msg { + err := m.client.SetRuleEnabled(r.ID, !r.Enabled) + return ruleWriteDoneMsg{action: "toggled", err: err} + } +} + +// enterRuleForm opens the add/edit form. r is nil to add a new rule. +func (m *model) enterRuleForm(r *rules.Rule) { + m.ruleForm = true + m.ruleField = fieldName + if r == nil { + m.ruleEditingID = 0 + m.ruleEnabled = true + m.ruleName.SetValue("") + m.ruleMatch.SetValue("") + m.ruleReplace.SetValue("") + m.ruleScope = "request" + m.ruleRegex = false + } else { + m.ruleEditingID = r.ID + m.ruleEnabled = r.Enabled + m.ruleName.SetValue(r.Name) + m.ruleMatch.SetValue(r.Match) + m.ruleReplace.SetValue(r.Replace) + m.ruleScope = r.Scope + m.ruleRegex = r.IsRegex + } + m.ruleName.Focus() + m.ruleMatch.Blur() + m.ruleReplace.Blur() +} + +func (m *model) ruleFromForm() rules.Rule { + return rules.Rule{ + ID: m.ruleEditingID, + Enabled: m.ruleEnabled, + Name: m.ruleName.Value(), + Scope: m.ruleScope, + Part: "header", + Match: m.ruleMatch.Value(), + Replace: m.ruleReplace.Value(), + IsRegex: m.ruleRegex, + } +} + +// focusRuleField moves input focus to m.ruleField, blurring the others. +func (m *model) focusRuleField() { + m.ruleName.Blur() + m.ruleMatch.Blur() + m.ruleReplace.Blur() + switch m.ruleField { + case fieldName: + m.ruleName.Focus() + case fieldMatch: + m.ruleMatch.Focus() + case fieldReplace: + m.ruleReplace.Focus() + } +} + func (m *model) Init() tea.Cmd { return tea.Batch(m.loadList, m.waitForEntry) } @@ -245,6 +393,13 @@ func (m *model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { m.reqArea.SetWidth(msg.Width) m.reqArea.SetHeight(reqHeight) m.respView = viewport.New(msg.Width, msg.Height-6-reqHeight) + + m.rulesTable.SetWidth(msg.Width) + m.rulesTable.SetHeight(msg.Height - 5) + formWidth := msg.Width - 12 + m.ruleName.Width = formWidth + m.ruleMatch.Width = formWidth + m.ruleReplace.Width = formWidth return m, nil case listLoadedMsg: @@ -300,6 +455,24 @@ func (m *model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { m.respView.GotoTop() return m, nil + case rulesLoadedMsg: + if msg.err != nil { + m.statusMsg = "rules error: " + msg.err.Error() + return m, nil + } + m.ruleRows = msg.rules + m.rulesTable.SetRows(rulesRowsFor(m.ruleRows)) + return m, nil + + case ruleWriteDoneMsg: + if msg.err != nil { + m.statusMsg = msg.action + " error: " + msg.err.Error() + return m, nil + } + m.ruleForm = false + m.statusMsg = "rule " + msg.action + return m, m.loadRules + case tea.KeyMsg: switch m.mode { case viewList: @@ -341,6 +514,10 @@ func (m *model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { m.searchInput.SetValue(m.query) m.searchInput.CursorEnd() return m, m.searchInput.Focus() + case "m": + m.mode = viewRules + m.statusMsg = "" + return m, m.loadRules case "esc": if m.query != "" { m.query = "" @@ -413,6 +590,78 @@ func (m *model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { m.respView, cmd = m.respView.Update(msg) } return m, cmd + + case viewRules: + if m.ruleForm { + switch msg.String() { + case "esc": + m.ruleForm = false + return m, nil + case "ctrl+c": + return m, tea.Quit + case "ctrl+s": + return m, m.saveRule(m.ruleFromForm()) + case "tab": + m.ruleField = (m.ruleField + 1) % 5 + m.focusRuleField() + return m, nil + case "shift+tab": + m.ruleField = (m.ruleField + 4) % 5 + m.focusRuleField() + return m, nil + } + if m.ruleField == fieldScope || m.ruleField == fieldRegex { + switch msg.String() { + case "left", "right", "enter", " ": + if m.ruleField == fieldScope { + if m.ruleScope == "request" { + m.ruleScope = "response" + } else { + m.ruleScope = "request" + } + } else { + m.ruleRegex = !m.ruleRegex + } + return m, nil + } + } + var cmd tea.Cmd + switch m.ruleField { + case fieldName: + m.ruleName, cmd = m.ruleName.Update(msg) + case fieldMatch: + m.ruleMatch, cmd = m.ruleMatch.Update(msg) + case fieldReplace: + m.ruleReplace, cmd = m.ruleReplace.Update(msg) + } + return m, cmd + } + + switch msg.String() { + case "q", "esc": + m.mode = viewList + return m, nil + case "ctrl+c": + return m, tea.Quit + case "a": + m.enterRuleForm(nil) + return m, nil + case "enter", "e": + if row := m.rulesTable.Cursor(); row >= 0 && row < len(m.ruleRows) { + sel := m.ruleRows[row] + m.enterRuleForm(&sel) + } + return m, nil + case "d": + m.statusMsg = "" + return m, m.deleteSelectedRule() + case " ": + m.statusMsg = "" + return m, m.toggleSelectedRule() + } + var cmd tea.Cmd + m.rulesTable, cmd = m.rulesTable.Update(msg) + return m, cmd } } return m, nil @@ -427,6 +676,11 @@ func (m *model) View() string { return m.detailView() case viewRepeater: return m.repeaterView() + case viewRules: + if m.ruleForm { + return m.ruleFormView() + } + return m.rulesView() default: return m.listView() } @@ -458,9 +712,9 @@ func (m *model) listView() string { b.WriteString(statusStyle.Render(m.statusMsg)) b.WriteString("\n") } - help := "↑/↓ navigate · enter view · r repeater · / search · q quit" + help := "↑/↓ navigate · enter view · r repeater · / search · m rules · q quit" if m.query != "" { - help = "↑/↓ navigate · enter view · r repeater · / search · esc clear filter · q quit" + help = "↑/↓ navigate · enter view · r repeater · / search · m rules · esc clear filter · q quit" } b.WriteString(helpStyle.Render(help)) return b.String() @@ -523,6 +777,75 @@ func (m *model) repeaterView() string { return b.String() } +func (m *model) rulesView() string { + var b strings.Builder + b.WriteString(titleStyle.Render(fmt.Sprintf(" match/replace rules (%d) - headers only for now ", len(m.ruleRows)))) + b.WriteString("\n") + b.WriteString(m.rulesTable.View()) + b.WriteString("\n") + if m.statusMsg != "" { + b.WriteString(statusStyle.Render(m.statusMsg)) + b.WriteString("\n") + } + b.WriteString(helpStyle.Render("a add · enter/e edit · d delete · space toggle · esc back · q quit")) + return b.String() +} + +func (m *model) ruleFormView() string { + var b strings.Builder + title := " add rule " + if m.ruleEditingID != 0 { + title = fmt.Sprintf(" edit rule #%d ", m.ruleEditingID) + } + b.WriteString(titleStyle.Render(title)) + b.WriteString("\n\n") + + label := func(field ruleField, text string) string { + if m.ruleField == field { + return tabActive.Render(text) + } + return tabInactive.Render(text) + } + + b.WriteString(label(fieldName, "Name") + "\n") + b.WriteString(m.ruleName.View() + "\n\n") + b.WriteString(label(fieldMatch, "Match") + "\n") + b.WriteString(m.ruleMatch.View() + "\n\n") + b.WriteString(label(fieldReplace, "Replace") + "\n") + b.WriteString(m.ruleReplace.View() + "\n\n") + + scopeText := fmt.Sprintf("Scope: %s (◀▶ to change)", m.ruleScope) + b.WriteString(label(fieldScope, scopeText) + "\n\n") + regexText := "Regex: off (◀▶ to change)" + if m.ruleRegex { + regexText = "Regex: on (◀▶ to change)" + } + b.WriteString(label(fieldRegex, regexText) + "\n\n") + + if m.statusMsg != "" { + b.WriteString(statusStyle.Render(m.statusMsg)) + b.WriteString("\n") + } + b.WriteString(helpStyle.Render("tab/shift+tab move · ctrl+s save · esc cancel · ctrl+c quit")) + return b.String() +} + +func rulesRowsFor(rs []rules.Rule) []table.Row { + rows := make([]table.Row, len(rs)) + for i, r := range rs { + on := " " + if r.Enabled { + on = "✓" + } + regex := "" + if r.IsRegex { + regex = "yes" + } + rows[i] = table.Row{on, r.Name, r.Scope, r.Match, r.Replace, regex} + } + return rows +} + func exactSuffix(exact bool) string { if exact { return ", exact" diff --git a/internal/ipc/ipc.go b/internal/ipc/ipc.go index f890715..19c9c23 100644 --- a/internal/ipc/ipc.go +++ b/internal/ipc/ipc.go @@ -11,12 +11,13 @@ import ( "net" "sync" + "mitmux/internal/rules" "mitmux/internal/store" ) // Request is sent by a client to the daemon. type Request struct { - Type string `json:"type"` // "list", "get", "subscribe", or "repeat" + Type string `json:"type"` // "list", "get", "subscribe", "repeat", "rules_list", "rules_save", "rules_delete", or "rules_toggle" Limit int `json:"limit,omitempty"` BeforeID int64 `json:"before_id,omitempty"` ID int64 `json:"id,omitempty"` @@ -30,14 +31,22 @@ type Request struct { Scheme string `json:"scheme,omitempty"` Host string `json:"host,omitempty"` Raw []byte `json:"raw,omitempty"` + + // For "rules_save": add (Rule.ID == 0) or update (Rule.ID != 0) a + // match-and-replace rule. For "rules_delete"/"rules_toggle": RuleID + // (and RuleEnabled for toggle) identify the target. + Rule *rules.Rule `json:"rule,omitempty"` + RuleID int64 `json:"rule_id,omitempty"` + RuleEnabled bool `json:"rule_enabled,omitempty"` } // Response is sent by the daemon to a client. type Response struct { - Type string `json:"type"` // "list", "get", "new", "repeat", or "error" + Type string `json:"type"` // "list", "get", "new", "repeat", "rules", or "error" Entries []store.Summary `json:"entries,omitempty"` // for "list" Detail *EntryDetail `json:"detail,omitempty"` // for "get" and "repeat" New *store.Summary `json:"new,omitempty"` // for "new" (subscribe push) + Rules []rules.Rule `json:"rules,omitempty"` // for "rules" Error string `json:"error,omitempty"` } @@ -144,6 +153,58 @@ func (c *Client) Repeat(scheme, host string, raw []byte) (*EntryDetail, error) { return resp.Detail, nil } +// ListRules returns every match-and-replace rule. +func (c *Client) ListRules() ([]rules.Rule, error) { + c.mu.Lock() + defer c.mu.Unlock() + return c.rulesRoundTrip(Request{Type: "rules_list"}) +} + +// SaveRule adds r (if r.ID == 0) or updates the existing rule with that +// ID, and returns its ID. +func (c *Client) SaveRule(r rules.Rule) (int64, error) { + c.mu.Lock() + defer c.mu.Unlock() + saved, err := c.rulesRoundTrip(Request{Type: "rules_save", Rule: &r}) + if err != nil { + return 0, err + } + if len(saved) == 0 { + return 0, errors.New("rules_save: daemon returned no rule") + } + return saved[0].ID, nil +} + +// DeleteRule removes a rule. +func (c *Client) DeleteRule(id int64) error { + c.mu.Lock() + defer c.mu.Unlock() + _, err := c.rulesRoundTrip(Request{Type: "rules_delete", RuleID: id}) + return err +} + +// SetRuleEnabled toggles a rule without touching its other fields. +func (c *Client) SetRuleEnabled(id int64, enabled bool) error { + c.mu.Lock() + defer c.mu.Unlock() + _, err := c.rulesRoundTrip(Request{Type: "rules_toggle", RuleID: id, RuleEnabled: enabled}) + return err +} + +func (c *Client) rulesRoundTrip(req Request) ([]rules.Rule, error) { + if err := c.enc.Encode(req); err != nil { + return nil, err + } + var resp Response + if err := c.dec.Decode(&resp); err != nil { + return nil, err + } + if resp.Type == "error" { + return nil, errors.New(resp.Error) + } + return resp.Rules, nil +} + // Subscribe opens a dedicated connection that streams newly captured // history entries as they happen. The returned channel is closed when // the connection ends; call the returned close func to stop early. diff --git a/internal/ipc/server.go b/internal/ipc/server.go index 69044af..1fe8d8f 100644 --- a/internal/ipc/server.go +++ b/internal/ipc/server.go @@ -8,6 +8,7 @@ import ( "net" "sync" + "mitmux/internal/rules" "mitmux/internal/store" ) @@ -127,6 +128,46 @@ func (s *Server) handleConn(conn net.Conn) { } enc.Encode(Response{Type: "repeat", Detail: detailFromEntry(e)}) + case "rules_list": + rs, err := s.db.ListRules() + if err != nil { + enc.Encode(Response{Type: "error", Error: err.Error()}) + continue + } + enc.Encode(Response{Type: "rules", Rules: rs}) + + case "rules_save": + if req.Rule == nil { + enc.Encode(Response{Type: "error", Error: "rules_save: missing rule"}) + continue + } + r := *req.Rule + var err error + if r.ID == 0 { + r.ID, err = s.db.AddRule(r) + } else { + err = s.db.UpdateRule(r) + } + if err != nil { + enc.Encode(Response{Type: "error", Error: err.Error()}) + continue + } + enc.Encode(Response{Type: "rules", Rules: []rules.Rule{r}}) + + case "rules_delete": + if err := s.db.DeleteRule(req.RuleID); err != nil { + enc.Encode(Response{Type: "error", Error: err.Error()}) + continue + } + enc.Encode(Response{Type: "rules"}) + + case "rules_toggle": + if err := s.db.SetRuleEnabled(req.RuleID, req.RuleEnabled); err != nil { + enc.Encode(Response{Type: "error", Error: err.Error()}) + continue + } + enc.Encode(Response{Type: "rules"}) + case "subscribe": sub := s.hub.subscribe() defer s.hub.unsubscribe(sub) diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index 014c601..3b2d50d 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -35,6 +35,7 @@ import ( "golang.org/x/net/http2" "mitmux/internal/ca" + "mitmux/internal/rules" "mitmux/internal/store" ) @@ -244,6 +245,17 @@ func (s *Server) forward(dial dialer, scheme, hostname string, w http.ResponseWr outReq.URL.Host = hostname outReq.RequestURI = "" stripHopByHop(outReq.Header) + // Header rules are applied to outReq only, after cloning and header + // stripping - history's request_raw keeps showing what the client + // actually sent (clientTee/reqBodyCap already capture from r, not + // outReq), while what actually reaches the upstream server reflects + // the rules. That split is deliberate: match-and-replace is a wire + // transform, not a rewrite of the audit trail. + if reqRules, err := s.enabledRules("request"); err != nil { + log.Printf("load request rules: %v", err) + } else if len(reqRules) > 0 { + outReq.Header = rules.ApplyHeaders(outReq.Header, reqRules) + } // Only needed when the client leg isn't tee-captured (HTTP/2): tee // the body as it streams through so the reconstructed capture isn't @@ -297,6 +309,17 @@ func (s *Server) forward(dial dialer, scheme, hostname string, w http.ResponseWr resp.Body = io.NopCloser(respBodyCap) } + // Same split as the request side: response_raw keeps reflecting what + // the origin server actually sent (captured below, from upstreamTee + // or respBodyCap, both already wired to resp.Body independent of + // resp.Header), while the client actually receives the rule-modified + // headers. + if respRules, err := s.enabledRules("response"); err != nil { + log.Printf("load response rules: %v", err) + } else if len(respRules) > 0 { + resp.Header = rules.ApplyHeaders(resp.Header, respRules) + } + stripHopByHop(resp.Header) for k, vv := range resp.Header { for _, v := range vv { @@ -317,6 +340,17 @@ func (s *Server) forward(dial dialer, scheme, hostname string, w http.ResponseWr s.record(started, duration, scheme, hostname, r, reqRaw, reqExact, respRaw, respExact, resp.StatusCode, "") } +// enabledRules fetches the current enabled match-and-replace rules for +// scope ("request" or "response") fresh from the store on every call - +// simple and always current, and cheap enough (a local, in-process +// SQLite query) not to bother caching for how this is actually used. +func (s *Server) enabledRules(scope string) ([]rules.Rule, error) { + if s.store == nil { + return nil, nil + } + return s.store.EnabledRules(scope) +} + // record stores one history entry and notifies OnEntry. func (s *Server) record(started time.Time, duration time.Duration, scheme, host string, r *http.Request, reqRaw []byte, reqExact bool, respRaw []byte, respExact bool, status int, errMsg string) { diff --git a/internal/rules/rules.go b/internal/rules/rules.go new file mode 100644 index 0000000..1f50ec4 --- /dev/null +++ b/internal/rules/rules.go @@ -0,0 +1,97 @@ +// Package rules implements match-and-replace: user-defined rules that +// rewrite request/response headers as they pass through the proxy. +// Deliberately headers-only for now - see ApplyHeaders for why bodies +// are a separate, harder problem (noted as a follow-up in PLAN.md). +package rules + +import ( + "bufio" + "net/http" + "net/textproto" + "regexp" + "sort" + "strings" +) + +// Rule is one match-and-replace rule. +type Rule struct { + ID int64 + Enabled bool + Name string + Scope string // "request" or "response" + Part string // "header" (only part supported so far) + Match string + Replace string + IsRegex bool + // Position orders rule application (ascending) when several rules + // could touch the same text. + Position int +} + +// ApplyHeaders rewrites h in place by serializing it to a raw +// "Name: value\r\n" block, running every enabled rule with Part=="header" +// over that text (in Position order), and reparsing the result. Working +// on the raw text rather than per-value substitution is what lets a rule +// add or remove a header entirely, not just rewrite an existing value - +// matching how Burp's header match/replace works. If a rule's output +// doesn't parse back as valid headers, ApplyHeaders returns h unchanged +// rather than risk sending something corrupted. +func ApplyHeaders(h http.Header, rs []Rule) http.Header { + keys := make([]string, 0, len(h)) + for k := range h { + keys = append(keys, k) + } + sort.Strings(keys) + + var block strings.Builder + for _, k := range keys { + for _, v := range h[k] { + block.WriteString(k) + block.WriteString(": ") + block.WriteString(v) + block.WriteString("\r\n") + } + } + text := block.String() + + changed := false + for _, r := range sortedByPosition(rs) { + if !r.Enabled || r.Part != "header" { + continue + } + if next, ok := apply(text, r); ok { + text, changed = next, true + } + } + if !changed { + return h + } + + tp := textproto.NewReader(bufio.NewReader(strings.NewReader(text + "\r\n"))) + mh, err := tp.ReadMIMEHeader() + if err != nil { + return h + } + return http.Header(mh) +} + +func sortedByPosition(rs []Rule) []Rule { + out := make([]Rule, len(rs)) + copy(out, rs) + sort.SliceStable(out, func(i, j int) bool { return out[i].Position < out[j].Position }) + return out +} + +func apply(text string, r Rule) (string, bool) { + if r.IsRegex { + re, err := regexp.Compile(r.Match) + if err != nil { + return text, false + } + return re.ReplaceAllString(text, r.Replace), true + } + if r.Match == "" { + return text, false + } + return strings.ReplaceAll(text, r.Match, r.Replace), true +} diff --git a/internal/store/store.go b/internal/store/store.go index a3c6b03..211162e 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -12,6 +12,8 @@ import ( "time" _ "modernc.org/sqlite" + + "mitmux/internal/rules" ) const schema = ` @@ -36,6 +38,18 @@ CREATE VIRTUAL TABLE IF NOT EXISTS history_fts USING fts5( method, host, path, request_text, response_text, tokenize = 'unicode61 remove_diacritics 2' ); + +CREATE TABLE IF NOT EXISTS rules ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + enabled INTEGER NOT NULL DEFAULT 1, + name TEXT NOT NULL DEFAULT '', + scope TEXT NOT NULL, + part TEXT NOT NULL, + match TEXT NOT NULL, + replace TEXT NOT NULL, + is_regex INTEGER NOT NULL DEFAULT 0, + position INTEGER NOT NULL DEFAULT 0 +); ` // Store is a handle to the history database. Safe for concurrent use. @@ -312,6 +326,101 @@ func prepareFTSQuery(q string) string { return strings.Join(fields, " ") } +// ListRules returns every match-and-replace rule, ordered for application. +func (s *Store) ListRules() ([]rules.Rule, error) { + rows, err := s.db.Query( + `SELECT id, enabled, name, scope, part, match, replace, is_regex, position + FROM rules ORDER BY position, id`, + ) + if err != nil { + return nil, fmt.Errorf("list rules: %w", err) + } + defer rows.Close() + + var out []rules.Rule + for rows.Next() { + var r rules.Rule + var enabled, isRegex int + if err := rows.Scan(&r.ID, &enabled, &r.Name, &r.Scope, &r.Part, &r.Match, &r.Replace, &isRegex, &r.Position); err != nil { + return nil, fmt.Errorf("scan rule row: %w", err) + } + r.Enabled = enabled != 0 + r.IsRegex = isRegex != 0 + out = append(out, r) + } + return out, rows.Err() +} + +// EnabledRules returns enabled rules for scope ("request" or +// "response"), ordered for application. +func (s *Store) EnabledRules(scope string) ([]rules.Rule, error) { + rows, err := s.db.Query( + `SELECT id, enabled, name, scope, part, match, replace, is_regex, position + FROM rules WHERE enabled = 1 AND scope = ? ORDER BY position, id`, + scope, + ) + if err != nil { + return nil, fmt.Errorf("enabled rules: %w", err) + } + defer rows.Close() + + var out []rules.Rule + for rows.Next() { + var r rules.Rule + var enabled, isRegex int + if err := rows.Scan(&r.ID, &enabled, &r.Name, &r.Scope, &r.Part, &r.Match, &r.Replace, &isRegex, &r.Position); err != nil { + return nil, fmt.Errorf("scan rule row: %w", err) + } + r.Enabled = enabled != 0 + r.IsRegex = isRegex != 0 + out = append(out, r) + } + return out, rows.Err() +} + +// AddRule stores r and returns its assigned ID. +func (s *Store) AddRule(r rules.Rule) (int64, error) { + res, err := s.db.Exec( + `INSERT INTO rules (enabled, name, scope, part, match, replace, is_regex, position) + VALUES (?, ?, ?, ?, ?, ?, ?, ?)`, + boolToInt(r.Enabled), r.Name, r.Scope, r.Part, r.Match, r.Replace, boolToInt(r.IsRegex), r.Position, + ) + if err != nil { + return 0, fmt.Errorf("add rule: %w", err) + } + return res.LastInsertId() +} + +// UpdateRule replaces the stored rule with the same ID as r. +func (s *Store) UpdateRule(r rules.Rule) error { + _, err := s.db.Exec( + `UPDATE rules SET enabled = ?, name = ?, scope = ?, part = ?, match = ?, replace = ?, is_regex = ?, position = ? + WHERE id = ?`, + boolToInt(r.Enabled), r.Name, r.Scope, r.Part, r.Match, r.Replace, boolToInt(r.IsRegex), r.Position, r.ID, + ) + if err != nil { + return fmt.Errorf("update rule %d: %w", r.ID, err) + } + return nil +} + +// SetRuleEnabled toggles a rule without touching its other fields. +func (s *Store) SetRuleEnabled(id int64, enabled bool) error { + _, err := s.db.Exec(`UPDATE rules SET enabled = ? WHERE id = ?`, boolToInt(enabled), id) + if err != nil { + return fmt.Errorf("set rule %d enabled: %w", id, err) + } + return nil +} + +// DeleteRule removes a rule. +func (s *Store) DeleteRule(id int64) error { + if _, err := s.db.Exec(`DELETE FROM rules WHERE id = ?`, id); err != nil { + return fmt.Errorf("delete rule %d: %w", id, err) + } + return nil +} + func boolToInt(b bool) int { if b { return 1 |