First Commit
ci / go (push) Waiting to run
ci / go-db (agent) (push) Waiting to run
ci / go-db (config) (push) Waiting to run
ci / go-db (db) (push) Waiting to run
ci / go-db (evidence) (push) Waiting to run
ci / go-db (llmrec) (push) Waiting to run
ci / go-db (server) (push) Waiting to run
web / web (push) Waiting to run
docs / links (push) Canceled after 0s
detections / detections (push) Canceled after 0s

This commit is contained in:
dela
2026-10-09 08:38:16 +08:00
commit 0335d572de
756 changed files with 201663 additions and 0 deletions
+178
View File
@@ -0,0 +1,178 @@
package db
import "testing"
// TestActivityPageSessions covers the reverse-paginated, per-session history added
// for the SSE remediation: Main/Plan/Worker filtering, before-cursor paging without
// gaps/overlap, hasMore, and the task-level snapshot cursor. Mirrors docs §11.1.
func TestActivityPageSessions(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
expID, err := d.CreateExploration("test", "分页历史")
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM explorations WHERE id=$1`, expID)
es := d.Exploration(expID)
intentA, err := es.AddIntent(map[string]any{"summary": "intent A"}, 5, nil, "planner")
if err != nil {
t.Fatal(err)
}
intentB, err := es.AddIntent(map[string]any{"summary": "intent B"}, 5, nil, "planner")
if err != nil {
t.Fatal(err)
}
// Interleave a mix of agents so a session filter must actually discriminate.
// 25 main, 25 planner (Goal+Planner share worker=planner), 30 workerA, 5 workerB.
appendN := func(n int, a Activity) {
for range n {
if _, err := es.AppendActivity(a); err != nil {
t.Fatal(err)
}
}
}
// Interleaving order matters: emit round-robin-ish so ids of one session are
// scattered, proving the WHERE filter (not a contiguous range) is what selects.
for range 25 {
appendN(1, Activity{Worker: "mainagent", Kind: "text", Summary: "m"})
appendN(1, Activity{Worker: "planner", Kind: "text", Summary: "p"})
appendN(1, Activity{NodeID: &intentA, Worker: "work#1", Kind: "text", Summary: "a"})
}
appendN(5, Activity{NodeID: &intentA, Worker: "work#1", Kind: "text", Summary: "a2"}) // workerA → 30 total
appendN(5, Activity{NodeID: &intentB, Worker: "work#2", Kind: "text", Summary: "b"})
// snapshot cursor = max id across the whole task.
snap, err := es.ActivityMaxID()
if err != nil {
t.Fatal(err)
}
// Helper: page through a whole session backward and assert coverage.
collect := func(f ActivitySessionFilter, pageSize int) []Activity {
var all []Activity
before := int64(0)
seen := map[int64]bool{}
for {
items, hasMore, err := es.ActivityPage(f, before, pageSize)
if err != nil {
t.Fatal(err)
}
// ascending order within a page
for i := 1; i < len(items); i++ {
if items[i-1].ID >= items[i].ID {
t.Fatalf("page not ascending: %d >= %d", items[i-1].ID, items[i].ID)
}
}
// no overlap across pages
for _, a := range items {
if seen[a.ID] {
t.Fatalf("duplicate id %d across pages", a.ID)
}
seen[a.ID] = true
}
all = append([]Activity{}, append(items, all...)...) // prepend older page
if !hasMore || len(items) == 0 {
break
}
before = items[0].ID
}
return all
}
main := collect(ActivitySessionFilter{Worker: "mainagent"}, 10)
if len(main) != 25 {
t.Fatalf("main count = %d, want 25", len(main))
}
plan := collect(ActivitySessionFilter{Worker: "planner"}, 7)
if len(plan) != 25 {
t.Fatalf("plan count = %d, want 25", len(plan))
}
wa := collect(ActivitySessionFilter{NodeID: &intentA}, 8)
if len(wa) != 30 {
t.Fatalf("workerA count = %d, want 30", len(wa))
}
wb := collect(ActivitySessionFilter{NodeID: &intentB}, 8)
if len(wb) != 5 {
t.Fatalf("workerB count = %d, want 5", len(wb))
}
// full ascending order across the reconstructed session
for i := 1; i < len(wa); i++ {
if wa[i-1].ID >= wa[i].ID {
t.Fatalf("reconstructed session not ascending at %d", i)
}
}
// latest page (before=0) must include the session's newest record.
latest, hasMore, err := es.ActivityPage(ActivitySessionFilter{NodeID: &intentA}, 0, 10)
if err != nil {
t.Fatal(err)
}
if !hasMore {
t.Fatalf("workerA should have more than one page")
}
if latest[len(latest)-1].ID != wa[len(wa)-1].ID {
t.Fatalf("latest page missing newest record")
}
// snapshot cursor is the whole-task max, ≥ any session's max.
if snap < wa[len(wa)-1].ID {
t.Fatalf("snapshot %d < workerA max %d", snap, wa[len(wa)-1].ID)
}
}
// TestListByKindPage covers the paged worker(intent) list that lets the session list
// reach past the old fixed 300 cap (docs §8 / §11.1 item 10).
func TestListByKindPage(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
expID, err := d.CreateExploration("test", "意图分页")
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM explorations WHERE id=$1`, expID)
es := d.Exploration(expID)
const total = 25
for range total {
if _, err := es.AddIntent(map[string]any{"summary": "i"}, 1, nil, "planner"); err != nil {
t.Fatal(err)
}
}
// Page backward in chunks of 10; expect 10,10,5 and hasMore false on last.
seen := map[int64]bool{}
before := int64(0)
pages := 0
for {
items, hasMore, err := es.ListByKindPage(KindIntent, before, 10)
if err != nil {
t.Fatal(err)
}
pages++
for _, n := range items {
if seen[n.ID] {
t.Fatalf("dup intent %d across pages", n.ID)
}
seen[n.ID] = true
}
if len(items) == 0 || !hasMore {
break
}
// newest-first within a page → the oldest (smallest id) is last; page older
// history before it next.
before = items[len(items)-1].ID
}
if len(seen) != total {
t.Fatalf("paged intents = %d, want %d", len(seen), total)
}
if pages < 3 {
t.Fatalf("expected ≥3 pages for %d items at size 10, got %d", total, pages)
}
}
+637
View File
@@ -0,0 +1,637 @@
package db
import (
"fmt"
"strconv"
"strings"
"unicode"
)
// Expr is one leaf DSL clause.
type Expr struct {
Field string // empty = bare-text full-text search
Op string // "=", "==", "!=", ">", ">=", "<", "<="
Value string
}
// astNode is a node in the parsed DSL expression tree.
type astNode struct {
kind string // "and", "or", "leaf"
children []*astNode
expr *Expr // only for "leaf"
}
func andNode(cs []*astNode) *astNode { return &astNode{kind: "and", children: cs} }
func orNode(cs []*astNode) *astNode { return &astNode{kind: "or", children: cs} }
func leafNode(e Expr) *astNode { return &astNode{kind: "leaf", expr: &e} }
// knownStringFields maps DSL field name → SQL column name.
// NOTE: "type" is intentionally excluded — it is a separate parameter, not a DSL field.
var knownStringFields = map[string]string{
"domain": "domain",
"root_domain": "root_domain",
"ip": "ip",
"url": "url",
"page_title": "page_title",
"title": "page_title",
"icp": "icp",
"service_name": "service_name",
"app_name": "app_name",
"bundle_id": "bundle_id",
"category": "category",
"app_icp": "app_icp",
"method": "method",
"service_type": "service_type",
"record_type": "record_type",
}
// knownArrayFields maps DSL field name → SQL column name (array).
var knownArrayFields = map[string]string{
"technology": "technologies",
"technologies": "technologies",
"tech": "technologies",
}
// knownNumericFields maps DSL field name → SQL column name (integer).
var knownNumericFields = map[string]string{
"port": "port",
"status_code": "status_code",
"status": "status_code",
}
func isKnownField(f string) bool {
f = strings.ToLower(f)
_, s := knownStringFields[f]
_, a := knownArrayFields[f]
_, n := knownNumericFields[f]
return s || a || n || f == "company_id" || f == "task_id"
}
// ── tokeniser ────────────────────────────────────────────────────────────────
const (
tkField = "FIELD"
tkBare = "BARE"
tkAnd = "AND"
tkOr = "OR"
tkLP = "LPAREN"
tkRP = "RPAREN"
tkEOF = "EOF"
)
type tok struct {
kind string
expr *Expr // set for tkField and tkBare
}
func tokenize(s string) ([]tok, error) {
var tokens []tok
i := 0
for i < len(s) {
for i < len(s) && unicode.IsSpace(rune(s[i])) {
i++
}
if i >= len(s) {
break
}
switch s[i] {
case '(':
tokens = append(tokens, tok{kind: tkLP})
i++
case ')':
tokens = append(tokens, tok{kind: tkRP})
i++
default:
if expr, end, ok := tryParseFieldExpr(s, i); ok {
tokens = append(tokens, tok{kind: tkField, expr: &expr})
i = end
continue
}
word, end := readToken(s, i)
if word == "" {
i++
continue
}
switch strings.ToUpper(word) {
case "AND":
tokens = append(tokens, tok{kind: tkAnd})
case "OR":
tokens = append(tokens, tok{kind: tkOr})
default:
e := Expr{Field: "", Op: "=", Value: word}
tokens = append(tokens, tok{kind: tkBare, expr: &e})
}
i = end
}
}
tokens = append(tokens, tok{kind: tkEOF})
return tokens, nil
}
// ── parser ───────────────────────────────────────────────────────────────────
//
// Grammar (AND binds tighter than OR):
// expr = or_expr
// or_expr = and_expr (OR and_expr)*
// and_expr = atom (AND atom)*
// atom = FIELD | BARE | '(' expr ')'
type dslParser struct {
tokens []tok
pos int
}
func (p *dslParser) peek() tok {
if p.pos >= len(p.tokens) {
return tok{kind: tkEOF}
}
return p.tokens[p.pos]
}
func (p *dslParser) consume() tok {
t := p.peek()
p.pos++
return t
}
func (p *dslParser) parseOr() (*astNode, error) {
left, err := p.parseAnd()
if err != nil {
return nil, err
}
children := []*astNode{left}
for p.peek().kind == tkOr {
p.consume()
right, err := p.parseAnd()
if err != nil {
return nil, err
}
children = append(children, right)
}
if len(children) == 1 {
return children[0], nil
}
return orNode(children), nil
}
func (p *dslParser) parseAnd() (*astNode, error) {
left, err := p.parseAtom()
if err != nil {
return nil, err
}
children := []*astNode{left}
for p.peek().kind == tkAnd {
p.consume()
right, err := p.parseAtom()
if err != nil {
return nil, err
}
children = append(children, right)
}
if len(children) == 1 {
return children[0], nil
}
return andNode(children), nil
}
func (p *dslParser) parseAtom() (*astNode, error) {
t := p.peek()
switch t.kind {
case tkField, tkBare:
p.consume()
return leafNode(*t.expr), nil
case tkLP:
p.consume()
node, err := p.parseOr()
if err != nil {
return nil, err
}
if p.peek().kind != tkRP {
return nil, fmt.Errorf("DSL 语法错误:缺少右括号 ')'")
}
p.consume()
return node, nil
case tkEOF:
return nil, fmt.Errorf("DSL 语法错误:表达式不完整")
default:
return nil, fmt.Errorf("DSL 语法错误:意外的 token '%s'", t.kind)
}
}
// ParseDSL parses a DSL query string into an expression tree.
//
// Syntax:
//
// field=value fuzzy match (ILIKE '%value%')
// field==value exact match
// field!=value exclude fuzzy
// port>8080 numeric comparison
// bare word full-text fuzzy across all main text fields
//
// Operators: AND OR (case-insensitive), parentheses for grouping.
// AND binds tighter than OR.
func ParseDSL(s string) (*astNode, error) {
if strings.TrimSpace(s) == "" {
return nil, nil
}
tokens, err := tokenize(s)
if err != nil {
return nil, err
}
p := &dslParser{tokens: tokens}
node, err := p.parseOr()
if err != nil {
return nil, err
}
if p.peek().kind != tkEOF {
return nil, fmt.Errorf("DSL 语法错误:意外的内容 '%s'", p.peek().kind)
}
return node, nil
}
// ── SQL builder ──────────────────────────────────────────────────────────────
// fullTextCols are searched for bare-text tokens.
var fullTextCols = []string{
"domain", "root_domain", "ip", "url", "page_title",
"icp", "service_name", "app_name", "app_description",
}
type whereBuilder struct {
args []any
base int // placeholders are numbered base+1, base+2, …; 0 = the usual $1, $2, …
}
func (b *whereBuilder) next(v any) string {
b.args = append(b.args, v)
return fmt.Sprintf("$%d", b.base+len(b.args))
}
func (b *whereBuilder) build(node *astNode) (string, error) {
switch node.kind {
case "and":
parts := make([]string, 0, len(node.children))
for _, child := range node.children {
clause, err := b.build(child)
if err != nil {
return "", err
}
parts = append(parts, "("+clause+")")
}
return strings.Join(parts, " AND "), nil
case "or":
parts := make([]string, 0, len(node.children))
for _, child := range node.children {
clause, err := b.build(child)
if err != nil {
return "", err
}
parts = append(parts, "("+clause+")")
}
return strings.Join(parts, " OR "), nil
case "leaf":
return b.buildLeaf(*node.expr)
}
return "", fmt.Errorf("unknown node kind: %s", node.kind)
}
func (b *whereBuilder) buildLeaf(e Expr) (string, error) {
f := strings.ToLower(e.Field)
// bare-text: OR across all text fields + arrays
if f == "" {
p := b.next("%" + e.Value + "%")
var parts []string
for _, col := range fullTextCols {
parts = append(parts, col+" ILIKE "+p)
}
parts = append(parts,
"EXISTS (SELECT 1 FROM unnest(technologies) t(v) WHERE v ILIKE "+p+")",
"EXISTS (SELECT 1 FROM unnest(bound_domains) t(v) WHERE v ILIKE "+p+")",
)
return "(" + strings.Join(parts, " OR ") + ")", nil
}
// task_id: $N = ANY(task_ids)
if f == "task_id" {
n, err := strconv.ParseInt(e.Value, 10, 64)
if err != nil {
return "", fmt.Errorf("task_id 需要整数值: %s", e.Value)
}
return b.next(n) + " = ANY(task_ids)", nil
}
// company_id: exact integer
if f == "company_id" {
n, err := strconv.ParseInt(e.Value, 10, 64)
if err != nil {
return "", fmt.Errorf("company_id 需要整数值: %s", e.Value)
}
return "company_id = " + b.next(n), nil
}
// numeric fields
if col, ok := knownNumericFields[f]; ok {
n, err := strconv.Atoi(e.Value)
if err != nil {
return "", fmt.Errorf("字段 %s 需要整数值: %s", f, e.Value)
}
op := e.Op
if op == "==" {
op = "="
}
if op != "=" && op != "!=" && op != ">" && op != ">=" && op != "<" && op != "<=" {
return "", fmt.Errorf("字段 %s 不支持运算符 %s", f, e.Op)
}
return fmt.Sprintf("%s %s %s", col, op, b.next(n)), nil
}
// array fields
if col, ok := knownArrayFields[f]; ok {
switch e.Op {
case "==":
return b.next(e.Value) + " = ANY(" + col + ")", nil
case "!=":
return "NOT (" + b.next(e.Value) + " = ANY(" + col + "))", nil
case "=":
p := b.next("%" + e.Value + "%")
return "EXISTS (SELECT 1 FROM unnest(" + col + ") t(v) WHERE v ILIKE " + p + ")", nil
default:
return "", fmt.Errorf("数组字段 %s 不支持运算符 %s", f, e.Op)
}
}
// string fields
if col, ok := knownStringFields[f]; ok {
switch e.Op {
case "=":
return col + " ILIKE " + b.next("%"+e.Value+"%"), nil
case "==":
return col + " = " + b.next(e.Value), nil
case "!=":
return col + " NOT ILIKE " + b.next("%"+e.Value+"%"), nil
default:
return "", fmt.Errorf("字符串字段 %s 不支持运算符 %s", f, e.Op)
}
}
return "", fmt.Errorf("未知字段: %s", f)
}
func buildDSLWhere(node *astNode) (string, []any, error) {
return buildDSLWhereBase(node, 0)
}
// buildDSLWhereBase is buildDSLWhere with a placeholder offset: emitted args are
// numbered base+1 onward, leaving $1..$base free for the caller (e.g. a scope CTE
// that reserves $1 for the task id).
func buildDSLWhereBase(node *astNode, base int) (string, []any, error) {
if node == nil {
return "1=1", nil, nil
}
b := &whereBuilder{base: base}
clause, err := b.build(node)
if err != nil {
return "", nil, err
}
return clause, b.args, nil
}
// ── helpers (shared with parser) ─────────────────────────────────────────────
// tryParseFieldExpr tries to parse "field op value" at pos.
func tryParseFieldExpr(s string, pos int) (Expr, int, bool) {
i := pos
if i >= len(s) || !isIdentStart(s[i]) {
return Expr{}, pos, false
}
for i < len(s) && isIdentChar(s[i]) {
i++
}
field := strings.ToLower(s[pos:i])
if !isKnownField(field) {
return Expr{}, pos, false
}
if i >= len(s) {
return Expr{}, pos, false
}
var op string
switch {
case i+1 < len(s) && (s[i] == '=' || s[i] == '!' || s[i] == '>' || s[i] == '<') && s[i+1] == '=':
op = s[i : i+2]
i += 2
case s[i] == '>' || s[i] == '<' || s[i] == '=':
op = string(s[i])
i++
default:
return Expr{}, pos, false
}
value, end := readToken(s, i)
if end == i {
return Expr{}, pos, false
}
return Expr{Field: field, Op: op, Value: value}, end, true
}
// readToken reads a quoted or unquoted token starting at pos.
func readToken(s string, pos int) (string, int) {
if pos >= len(s) {
return "", pos
}
if s[pos] == '"' {
i := pos + 1
for i < len(s) && s[i] != '"' {
i++
}
val := s[pos+1 : i]
if i < len(s) {
i++
}
return val, i
}
i := pos
for i < len(s) && !unicode.IsSpace(rune(s[i])) && s[i] != '(' && s[i] != ')' {
i++
}
return s[pos:i], i
}
func isIdentStart(c byte) bool {
return c == '_' || (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z')
}
func isIdentChar(c byte) bool {
return isIdentStart(c) || (c >= '0' && c <= '9')
}
// ── QueryDSL ─────────────────────────────────────────────────────────────────
// ValidateDSL parses and compiles a DSL expression without touching the
// database. HTTP callers use it to distinguish client syntax errors from query
// failures, which must remain server errors.
func ValidateDSL(dsl string) error {
node, err := ParseDSL(dsl)
if err != nil {
return err
}
_, _, err = buildDSLWhere(node)
return err
}
// CountDSL returns the total number of assets matching a DSL expression (and optional
// type), for server-side pagination — same WHERE as QueryDSL, without LIMIT/OFFSET.
// taskID > 0 scopes the count to assets attached to that task.
func (s *AssetStore) CountDSL(dsl, typ string, taskID int64) (int, error) {
node, err := ParseDSL(dsl)
if err != nil {
return 0, err
}
where, args, err := buildDSLWhere(node)
if err != nil {
return 0, err
}
if typ != "" {
args = append(args, typ)
where += fmt.Sprintf(" AND type = $%d", len(args))
}
if taskID > 0 {
args = append(args, taskID)
where += fmt.Sprintf(" AND $%d = ANY(task_ids)", len(args))
}
var n int
err = s.db.QueryRow("SELECT count(*) FROM assets WHERE "+where, args...).Scan(&n)
return n, err
}
// QueryDSL executes a DSL query string against the asset store.
// typ is an optional asset type filter applied independently of the DSL expression.
// taskID > 0 scopes results to assets attached to that task and hydrates each
// row's per-task source metadata (as QueryByTask does).
func (s *AssetStore) QueryDSL(dsl, typ string, taskID int64, limit, offset int) ([]*Asset, error) {
if limit <= 0 {
limit = 50
}
if offset < 0 {
offset = 0
}
node, err := ParseDSL(dsl)
if err != nil {
return nil, err
}
where, args, err := buildDSLWhere(node)
if err != nil {
return nil, err
}
if typ != "" {
args = append(args, typ)
where += fmt.Sprintf(" AND type = $%d", len(args))
}
if taskID > 0 {
args = append(args, taskID)
where += fmt.Sprintf(" AND $%d = ANY(task_ids)", len(args))
}
args = append(args, limit, offset)
q := assetSelectCols + " WHERE " + where +
fmt.Sprintf(" ORDER BY last_seen DESC, id DESC LIMIT $%d OFFSET $%d", len(args)-1, len(args))
rows, err := s.db.Query(q, args...)
if err != nil {
return nil, err
}
defer rows.Close()
assets, err := scanAssets(rows)
if err != nil {
return nil, err
}
if taskID > 0 {
if err := s.hydrateTaskAssetSources(taskID, assets); err != nil {
return nil, err
}
}
return assets, nil
}
// QueryDSLInScope is QueryDSL restricted to assets that BELONG to taskID's (and its
// direct source tasks') declared scope — membership, not literal value: a
// root_domain scope returns every subdomain / service / endpoint under it. This is
// the agent-facing list_assets path, so an agent queries the task's relevant assets
// instead of the whole shared库. taskID<=0 (non-task contexts: Auto / pentest / chat)
// has no scope to honor and falls back to the plain global QueryDSL. Rows carry the
// same per-task source metadata as QueryByTask.
func (s *AssetStore) QueryDSLInScope(dsl, typ string, taskID int64, limit, offset int) ([]*Asset, error) {
if taskID <= 0 {
return s.QueryDSL(dsl, typ, 0, limit, offset)
}
if limit <= 0 {
limit = 50
}
if offset < 0 {
offset = 0
}
node, err := ParseDSL(dsl)
if err != nil {
return nil, err
}
// $1 is reserved for taskID (scopeTargetCTE); DSL placeholders start at $2.
where, dslArgs, err := buildDSLWhereBase(node, 1)
if err != nil {
return nil, err
}
args := []any{taskID}
args = append(args, dslArgs...)
if typ != "" {
args = append(args, typ)
where += fmt.Sprintf(" AND type = $%d", len(args))
}
where += " AND id IN (SELECT id FROM target)"
args = append(args, limit, offset)
q := `WITH ` + scopeTargetCTE + ` ` + assetSelectCols + " WHERE " + where +
fmt.Sprintf(" ORDER BY last_seen DESC, id DESC LIMIT $%d OFFSET $%d", len(args)-1, len(args))
rows, err := s.db.Query(q, args...)
if err != nil {
return nil, err
}
defer rows.Close()
assets, err := scanAssets(rows)
if err != nil {
return nil, err
}
if err := s.hydrateTaskAssetSources(taskID, assets); err != nil {
return nil, err
}
return assets, nil
}
// GetByIDsInScope is GetByIDs restricted to ids that BELONG to taskID's (and its
// direct source tasks') declared scope, so an agent cannot reach out-of-scope
// assets by id. taskID<=0 (non-task contexts) falls back to the global GetByIDs.
// Out-of-scope ids are silently dropped from the result (not an error).
func (s *AssetStore) GetByIDsInScope(taskID int64, ids []int64) ([]*Asset, error) {
if taskID <= 0 {
return s.GetByIDs(ids)
}
if len(ids) == 0 {
return nil, nil
}
args := make([]any, 0, len(ids)+1)
args = append(args, taskID) // $1 reserved for scopeTargetCTE
placeholders := make([]string, len(ids))
for i, id := range ids {
placeholders[i] = fmt.Sprintf("$%d", i+2)
args = append(args, id)
}
q := `WITH ` + scopeTargetCTE + ` ` + assetSelectCols +
" WHERE id IN (" + strings.Join(placeholders, ",") + ")" +
" AND id IN (SELECT id FROM target) ORDER BY last_seen DESC, id DESC"
rows, err := s.db.Query(q, args...)
if err != nil {
return nil, err
}
defer rows.Close()
assets, err := scanAssets(rows)
if err != nil {
return nil, err
}
if err := s.hydrateTaskAssetSources(taskID, assets); err != nil {
return nil, err
}
return assets, nil
}
+80
View File
@@ -0,0 +1,80 @@
package db
import "time"
// AssetInterceptRule is one row of asset_intercept_rules — a global asset
// blocklist entry. Unlike intercept_rules (which matches tool name / input),
// these match the *target asset*: an exact/fuzzy domain·ip·url, or a CIDR range.
// This layer only stores rules; the matching/enforcement logic lives elsewhere.
type AssetInterceptRule struct {
ID int64 `json:"id"`
Enabled bool `json:"enabled"`
Kind string `json:"kind"` // exact_domain|exact_ip|exact_url|fuzzy_domain|fuzzy_ip|fuzzy_url|cidr
Pattern string `json:"pattern"`
Note string `json:"note"`
Builtin bool `json:"builtin"`
// Action 仅用于任务级规则:'block'=拦截 'allow'=允许(白名单)。
// 全局规则(asset_intercept_rules)不带此列,恒为空,视为拦截。
Action string `json:"action,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
const assetInterceptRuleCols = `id, enabled, kind, pattern, note, builtin, created_at, updated_at`
func scanAssetInterceptRule(row interface{ Scan(...any) error }) (AssetInterceptRule, error) {
var r AssetInterceptRule
err := row.Scan(&r.ID, &r.Enabled, &r.Kind, &r.Pattern, &r.Note, &r.Builtin, &r.CreatedAt, &r.UpdatedAt)
return r, err
}
// ListAssetInterceptRules returns all rules, built-ins first then newest first.
func (d *DB) ListAssetInterceptRules() ([]AssetInterceptRule, error) {
rows, err := d.Query(`SELECT ` + assetInterceptRuleCols + ` FROM asset_intercept_rules ORDER BY builtin DESC, id DESC`)
if err != nil {
return nil, err
}
defer rows.Close()
var out []AssetInterceptRule
for rows.Next() {
r, err := scanAssetInterceptRule(rows)
if err != nil {
return nil, err
}
out = append(out, r)
}
return out, rows.Err()
}
// CreateAssetInterceptRule inserts a new user rule (builtin is always false here).
func (d *DB) CreateAssetInterceptRule(kind, pattern, note string, enabled bool) (AssetInterceptRule, error) {
row := d.QueryRow(`
INSERT INTO asset_intercept_rules(enabled, kind, pattern, note, builtin)
VALUES ($1, $2, $3, $4, false)
RETURNING `+assetInterceptRuleCols,
enabled, kind, pattern, note)
return scanAssetInterceptRule(row)
}
// UpdateAssetInterceptRule replaces the editable fields of an existing rule.
func (d *DB) UpdateAssetInterceptRule(id int64, kind, pattern, note string, enabled bool) (AssetInterceptRule, error) {
row := d.QueryRow(`
UPDATE asset_intercept_rules
SET enabled=$2, kind=$3, pattern=$4, note=$5
WHERE id=$1
RETURNING `+assetInterceptRuleCols,
id, enabled, kind, pattern, note)
return scanAssetInterceptRule(row)
}
// DeleteAssetInterceptRule removes a rule (built-in rules are deletable too).
func (d *DB) DeleteAssetInterceptRule(id int64) error {
_, err := d.Exec(`DELETE FROM asset_intercept_rules WHERE id=$1`, id)
return err
}
// ToggleAssetInterceptRule flips the enabled state of a rule.
func (d *DB) ToggleAssetInterceptRule(id int64, enabled bool) error {
_, err := d.Exec(`UPDATE asset_intercept_rules SET enabled=$2 WHERE id=$1`, id, enabled)
return err
}
+250
View File
@@ -0,0 +1,250 @@
package db
import (
"fmt"
"net"
"net/url"
"strings"
)
// 资产拦截规则的匹配/执行层。asset_intercept.go 只负责规则存储,这里负责把
// 「目标资产」的域名/IP/URL 与启用中的规则做匹配。供 agent 工具(add_intent、
// insert_assets)在下发意图 / 插入资产前调用,命中则拒绝。
// AssetInterceptKindLabel 返回 kind 的中文标签,用于给 agent 的说明消息。
func AssetInterceptKindLabel(kind string) string {
switch kind {
case "exact_domain":
return "域名(全等)"
case "exact_ip":
return "IP(全等)"
case "exact_url":
return "URL(全等)"
case "fuzzy_domain":
return "域名(模糊)"
case "fuzzy_ip":
return "IP(模糊)"
case "fuzzy_url":
return "URL(模糊)"
case "cidr":
return "CIDR 网段"
}
return kind
}
// Reason 返回一条可读的命中原因,形如:命中资产拦截规则 [域名(模糊): .gov.cn](备注)。
func (r AssetInterceptRule) Reason() string {
s := fmt.Sprintf("命中资产拦截规则 [%s: %s]", AssetInterceptKindLabel(r.Kind), r.Pattern)
if note := strings.TrimSpace(r.Note); note != "" {
s += "(" + note + ")"
}
return s
}
// matchOne 判断单条启用规则是否命中给定的域名/IP/URL 候选串,返回命中的具体值。
func matchOne(r AssetInterceptRule, domains, ips, urls []string) (string, bool) {
p := strings.TrimSpace(r.Pattern)
if p == "" {
return "", false
}
switch r.Kind {
case "exact_domain":
for _, d := range domains {
if strings.EqualFold(strings.TrimSpace(d), p) {
return d, true
}
}
case "exact_ip":
for _, ip := range ips {
if strings.TrimSpace(ip) == p {
return ip, true
}
}
case "exact_url":
for _, u := range urls {
if strings.TrimSpace(u) == p {
return u, true
}
}
case "fuzzy_domain":
lp := strings.ToLower(p)
for _, d := range domains {
if d != "" && strings.Contains(strings.ToLower(d), lp) {
return d, true
}
}
case "fuzzy_ip":
for _, ip := range ips {
if ip != "" && strings.Contains(ip, p) {
return ip, true
}
}
case "fuzzy_url":
lp := strings.ToLower(p)
for _, u := range urls {
if u != "" && strings.Contains(strings.ToLower(u), lp) {
return u, true
}
}
case "cidr":
_, ipnet, err := net.ParseCIDR(p)
if err != nil {
return "", false
}
for _, ip := range ips {
if pip := net.ParseIP(strings.TrimSpace(ip)); pip != nil && ipnet.Contains(pip) {
return ip, true
}
}
}
return "", false
}
// MatchAssetInterceptRules 返回第一条命中给定 域名/IP/URL 候选串的启用规则,及命中的具体值。
// 供 insert_assets 用原始输入(尚未落库的 assetInputItem)匹配。
func MatchAssetInterceptRules(rules []AssetInterceptRule, domains, ips, urls []string) (AssetInterceptRule, string, bool) {
for _, r := range rules {
if !r.Enabled {
continue
}
if v, ok := matchOne(r, domains, ips, urls); ok {
return r, v, true
}
}
return AssetInterceptRule{}, "", false
}
// interceptCandidates 提取一个已落库资产用于拦截匹配的 域名/IP/URL 候选串。
// URL 的 host 会被拆出并归类,使「只带 URL」的服务类资产也能被 域名/IP 规则命中。
func (a *Asset) interceptCandidates() (domains, ips, urls []string) {
add := func(dst *[]string, s string) {
if s = strings.TrimSpace(s); s != "" {
*dst = append(*dst, s)
}
}
add(&domains, a.Domain)
add(&domains, a.RootDomain)
for _, d := range a.BoundDomains {
add(&domains, d)
}
add(&ips, a.IP)
add(&urls, a.URL)
if a.URL != "" {
if u, err := url.Parse(a.URL); err == nil {
if h := u.Hostname(); h != "" {
if net.ParseIP(h) != nil {
add(&ips, h)
} else {
add(&domains, h)
}
}
}
}
return domains, ips, urls
}
// InterceptLabel 返回资产的简短标识,用于给 agent 的说明消息。
func (a *Asset) InterceptLabel() string {
var target string
switch {
case a.Domain != "":
target = a.Domain
case a.URL != "":
target = a.URL
case a.IP != "":
target = a.IP
default:
target = fmt.Sprintf("#%d", a.ID)
}
return fmt.Sprintf("资产#%d[%s] %s", a.ID, a.Type, target)
}
// hasEnabledRule 判断规则集里是否存在任一启用规则。
func hasEnabledRule(rules []AssetInterceptRule) bool {
for _, r := range rules {
if r.Enabled {
return true
}
}
return false
}
// AssetGateDecision 是「先拦截后允许」闸门对一组候选串的判定结果。
type AssetGateDecision struct {
Allowed bool
Reason string // 被拒原因(不含资产标识);Allowed=true 时为空
}
// EvaluateAssetGate 执行任务级闸门判定:
// 1. 命中任一启用的 blockRules → 拒绝(拦截原因)。
// 2. 否则若 allowRules 存在启用项且都不命中 → 拒绝(不在允许范围)。
// 3. 否则放行。
//
// allowRules 为空/无启用项时,允许闸门不生效(即不启用白名单,全部放行),
// 避免「未配置允许规则」把所有资产挡掉。
func EvaluateAssetGate(blockRules, allowRules []AssetInterceptRule, domains, ips, urls []string) AssetGateDecision {
if rule, _, ok := MatchAssetInterceptRules(blockRules, domains, ips, urls); ok {
return AssetGateDecision{Allowed: false, Reason: rule.Reason()}
}
if hasEnabledRule(allowRules) {
if _, _, ok := MatchAssetInterceptRules(allowRules, domains, ips, urls); !ok {
return AssetGateDecision{Allowed: false, Reason: "不在任务允许(白名单)范围内,不允许测试"}
}
}
return AssetGateDecision{Allowed: true}
}
// AssetInterceptHit 描述一个被闸门拒绝的资产(拦截命中 或 不在允许范围)。
type AssetInterceptHit struct {
Asset *Asset
Reason string // 可读原因
}
// Describe 返回一条可读的说明:资产信息 + 原因。
func (h AssetInterceptHit) Describe() string {
return fmt.Sprintf("%s → %s", h.Asset.InterceptLabel(), h.Reason)
}
// ListAssetInterceptRules 是 *DB 同名方法的透传,让只持有 AssetStore 的调用方
// (如 agent 工具)也能读取规则。
func (s *AssetStore) ListAssetInterceptRules() ([]AssetInterceptRule, error) {
return s.db.ListAssetInterceptRules()
}
// CheckAssetsIntercept 按 id 载入资产,逐个执行「先拦截后允许」闸门判定,返回所有
// 被拒的资产。拦截规则 = 全局 ∪ 任务级 block;允许规则 = 任务级 allow(仅本任务)。
// 无 id 时快速返回。用全局 GetByIDs(不受任务范围过滤)以保证拦截不被 scope 削弱。
func (s *AssetStore) CheckAssetsIntercept(taskID int64, ids []int64) ([]AssetInterceptHit, error) {
if len(ids) == 0 {
return nil, nil
}
blockRules, err := s.db.ListAssetInterceptRules()
if err != nil {
return nil, err
}
var allowRules []AssetInterceptRule
if taskID > 0 {
tb, ta, err := s.TaskInterceptRulesSplit(taskID)
if err != nil {
return nil, err
}
blockRules = append(blockRules, tb...)
allowRules = ta
}
// 既无拦截规则、也无启用的允许规则 → 无需判定,全部放行。
if len(blockRules) == 0 && !hasEnabledRule(allowRules) {
return nil, nil
}
assets, err := s.GetByIDs(ids)
if err != nil {
return nil, err
}
var hits []AssetInterceptHit
for _, a := range assets {
domains, ips, urls := a.interceptCandidates()
if d := EvaluateAssetGate(blockRules, allowRules, domains, ips, urls); !d.Allowed {
hits = append(hits, AssetInterceptHit{Asset: a, Reason: d.Reason})
}
}
return hits, nil
}
+129
View File
@@ -0,0 +1,129 @@
package db
import "testing"
func rule(kind, pattern string, enabled bool) AssetInterceptRule {
return AssetInterceptRule{Kind: kind, Pattern: pattern, Enabled: enabled}
}
func TestMatchAssetInterceptRules(t *testing.T) {
cases := []struct {
name string
rules []AssetInterceptRule
domains []string
ips []string
urls []string
want bool
wantVal string
}{
{"内置模糊政府域名命中", []AssetInterceptRule{rule("fuzzy_domain", ".gov.cn", true)},
[]string{"www.beijing.gov.cn"}, nil, nil, true, "www.beijing.gov.cn"},
{"模糊教育域名命中", []AssetInterceptRule{rule("fuzzy_domain", ".edu", true)},
[]string{"mit.edu"}, nil, nil, true, "mit.edu"},
{"全等域名命中大小写不敏感", []AssetInterceptRule{rule("exact_domain", "Example.com", true)},
[]string{"example.com"}, nil, nil, true, "example.com"},
{"全等域名不命中子域", []AssetInterceptRule{rule("exact_domain", "example.com", true)},
[]string{"a.example.com"}, nil, nil, false, ""},
{"全等IP命中", []AssetInterceptRule{rule("exact_ip", "203.0.113.5", true)},
nil, []string{"203.0.113.5"}, nil, true, "203.0.113.5"},
{"模糊IP前缀命中", []AssetInterceptRule{rule("fuzzy_ip", "203.0.113.", true)},
nil, []string{"203.0.113.99"}, nil, true, "203.0.113.99"},
{"CIDR 命中", []AssetInterceptRule{rule("cidr", "192.168.0.0/16", true)},
nil, []string{"192.168.5.20"}, nil, true, "192.168.5.20"},
{"CIDR 不命中", []AssetInterceptRule{rule("cidr", "192.168.0.0/16", true)},
nil, []string{"10.0.0.1"}, nil, false, ""},
{"全等URL命中", []AssetInterceptRule{rule("exact_url", "https://a.gov.cn/login", true)},
nil, nil, []string{"https://a.gov.cn/login"}, true, "https://a.gov.cn/login"},
{"模糊URL命中路径", []AssetInterceptRule{rule("fuzzy_url", "/admin", true)},
nil, nil, []string{"https://x.com/admin/panel"}, true, "https://x.com/admin/panel"},
{"禁用规则不命中", []AssetInterceptRule{rule("fuzzy_domain", ".gov.cn", false)},
[]string{"www.gov.cn"}, nil, nil, false, ""},
{"无规则不命中", nil, []string{"www.gov.cn"}, nil, nil, false, ""},
{"空pattern不命中", []AssetInterceptRule{rule("fuzzy_domain", " ", true)},
[]string{"www.gov.cn"}, nil, nil, false, ""},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
r, val, ok := MatchAssetInterceptRules(c.rules, c.domains, c.ips, c.urls)
if ok != c.want {
t.Fatalf("命中 = %v, 期望 %v (rule=%+v)", ok, c.want, r)
}
if ok && val != c.wantVal {
t.Fatalf("命中值 = %q, 期望 %q", val, c.wantVal)
}
})
}
}
func TestEvaluateAssetGate(t *testing.T) {
block := []AssetInterceptRule{rule("fuzzy_domain", ".gov.cn", true)}
allow := []AssetInterceptRule{rule("fuzzy_domain", "example.com", true)}
// 1. 命中拦截规则 → 拒绝(拦截原因优先)。
if d := EvaluateAssetGate(block, allow, []string{"www.gov.cn"}, nil, nil); d.Allowed {
t.Fatal("命中拦截规则应被拒绝")
}
// 2. 未命中拦截、有允许规则但不命中 → 拒绝(不允许)。
d := EvaluateAssetGate(block, allow, []string{"foo.other.com"}, nil, nil)
if d.Allowed {
t.Fatal("有白名单且不命中应被拒绝")
}
if d.Reason == "" {
t.Fatal("拒绝应带原因")
}
// 3. 未命中拦截、命中允许规则 → 放行。
if d := EvaluateAssetGate(block, allow, []string{"api.example.com"}, nil, nil); !d.Allowed {
t.Fatal("命中白名单应放行")
}
// 4. 无允许规则(白名单未启用)→ 未命中拦截即放行。
if d := EvaluateAssetGate(block, nil, []string{"foo.other.com"}, nil, nil); !d.Allowed {
t.Fatal("无白名单时未命中拦截应放行")
}
// 5. 允许规则全部禁用 → 视为白名单未启用,放行。
disabledAllow := []AssetInterceptRule{rule("fuzzy_domain", "example.com", false)}
if d := EvaluateAssetGate(nil, disabledAllow, []string{"foo.other.com"}, nil, nil); !d.Allowed {
t.Fatal("白名单全禁用时应放行")
}
// 6. 拦截优先于允许:同一目标既命中拦截又命中允许 → 拒绝。
if d := EvaluateAssetGate(
[]AssetInterceptRule{rule("fuzzy_domain", ".gov.cn", true)},
[]AssetInterceptRule{rule("fuzzy_domain", ".gov.cn", true)},
[]string{"www.gov.cn"}, nil, nil,
); d.Allowed {
t.Fatal("拦截应优先于允许")
}
}
func TestAssetInterceptCandidates(t *testing.T) {
// 只带 URL 的服务资产:host 应被拆出并归入域名候选,从而被 fuzzy_domain 命中。
a := &Asset{Type: "service", URL: "https://portal.beijing.gov.cn:8443/app"}
domains, _, urls := a.interceptCandidates()
if len(urls) != 1 || urls[0] != a.URL {
t.Fatalf("urls = %v", urls)
}
found := false
for _, d := range domains {
if d == "portal.beijing.gov.cn" {
found = true
}
}
if !found {
t.Fatalf("URL host 未拆入域名候选: %v", domains)
}
r, _, ok := MatchAssetInterceptRules([]AssetInterceptRule{rule("fuzzy_domain", ".gov.cn", true)}, domains, nil, urls)
if !ok {
t.Fatalf("仅带 URL 的政府服务资产应被 fuzzy_domain 命中, rule=%+v", r)
}
// URL host 是 IP 时应归入 IP 候选,可被 CIDR 命中。
b := &Asset{Type: "service", URL: "http://10.1.2.3/x"}
_, ips, _ := b.interceptCandidates()
if r, _, ok := MatchAssetInterceptRules([]AssetInterceptRule{rule("cidr", "10.0.0.0/8", true)}, nil, ips, nil); !ok {
t.Fatalf("URL 中的 IP 应被 CIDR 命中, ips=%v rule=%+v", ips, r)
}
}
+1536
View File
File diff suppressed because it is too large Load Diff
+1126
View File
File diff suppressed because it is too large Load Diff
+113
View File
@@ -0,0 +1,113 @@
package db
import (
"context"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"strings"
)
// ChatMention is a lightweight search result. Details are read again on send.
type ChatMention struct {
Kind string `json:"kind"`
ID int64 `json:"id"`
Label string `json:"label"`
Description string `json:"description"`
}
type ChatMentionPage struct {
Items []ChatMention `json:"items"`
NextCursor string `json:"next_cursor,omitempty"`
}
var ErrInvalidChatMentionCursor = errors.New("分页位置无效,请重新搜索")
type chatMentionCursor struct {
ID int64 `json:"id"`
Kind string `json:"kind"`
Exact bool `json:"exact"`
Scope string `json:"scope"`
Query string `json:"query"`
}
func ValidChatMentionKind(kind string) bool {
switch kind {
case "finding", "company", "asset", "endpoint", "ip", "app", "root_domain", "subdomain", "service":
return true
}
return false
}
// SearchChatMentions searches the shared catalog, just like the asset/finding
// pages. Values remain SQL parameters; %, _ and backslash are literal search text.
func (d *DB) SearchChatMentions(ctx context.Context, kind, query string) ([]ChatMention, error) {
page, err := d.SearchChatMentionsPage(ctx, kind, query, "")
return page.Items, err
}
// SearchChatMentionsPage uses the last result's stable sort key rather than an
// offset, so loading later pages does not repeatedly skip all earlier results.
func (d *DB) SearchChatMentionsPage(ctx context.Context, kind, query, cursor string) (ChatMentionPage, error) {
page := ChatMentionPage{Items: make([]ChatMention, 0)}
if kind != "" && !ValidChatMentionKind(kind) {
return page, fmt.Errorf("不支持的引用类型")
}
var after chatMentionCursor
if cursor != "" {
if len(cursor) > 2048 {
return page, ErrInvalidChatMentionCursor
}
blob, err := base64.RawURLEncoding.DecodeString(cursor)
if err != nil || json.Unmarshal(blob, &after) != nil || after.ID <= 0 ||
!ValidChatMentionKind(after.Kind) || after.Scope != kind || after.Query != query {
return page, ErrInvalidChatMentionCursor
}
}
pattern := "%" + strings.NewReplacer(`\`, `\\`, "%", `\%`, "_", `\_`).Replace(query) + "%"
rows, err := d.QueryContext(ctx, `
SELECT kind, id, left(label, 160), left(description, 240) FROM (
(SELECT 'finding' AS kind, id, COALESCE(NULLIF(name,''), vulnclass) AS label,
concat_ws(' · ', severity, status, left(summary, 160)) AS description
FROM findings WHERE ($1='' OR $1='finding') AND
($2='' OR id::text=$2 OR concat_ws(' ',name,vulnclass,summary) ILIKE $3)
AND ($4::bigint=0 OR (id::text=$2)<$6 OR ((id::text=$2)=$6 AND (id<$4 OR (id=$4 AND 'finding'>$5))))
ORDER BY (id::text=$2) DESC, id DESC LIMIT 21)
UNION ALL
(SELECT 'company', id, name, nkey FROM companies
WHERE ($1='' OR $1='company') AND ($2='' OR id::text=$2 OR name ILIKE $3 OR nkey ILIKE $3)
AND ($4::bigint=0 OR (id::text=$2)<$6 OR ((id::text=$2)=$6 AND (id<$4 OR (id=$4 AND 'company'>$5))))
ORDER BY (id::text=$2) DESC, id DESC LIMIT 21)
UNION ALL
(SELECT type, id,
CASE WHEN type='endpoint' THEN concat_ws(' ',NULLIF(method,''),url)
ELSE COALESCE(NULLIF(app_name,''),NULLIF(url,''),NULLIF(domain,''),NULLIF(ip,''),NULLIF(bundle_id,''),'资产 #'||id::text) END,
concat_ws(' · ',type,NULLIF(page_title,''),NULLIF(service_name,''),NULLIF(bundle_id,''),NULLIF(ip,''),port::text)
FROM assets WHERE ($1='' OR $1='asset' OR type=$1) AND
($2='' OR id::text=$2 OR concat_ws(' ',domain,root_domain,ip,url,app_name,bundle_id,page_title,service_name,method) ILIKE $3)
AND ($4::bigint=0 OR (id::text=$2)<$6 OR ((id::text=$2)=$6 AND (id<$4 OR (id=$4 AND type>$5))))
ORDER BY (id::text=$2) DESC, id DESC LIMIT 21)
) matches ORDER BY (id::text=$2) DESC, id DESC, kind LIMIT 21`, kind, query, pattern, after.ID, after.Kind, after.Exact)
if err != nil {
return page, err
}
defer rows.Close()
for rows.Next() {
var item ChatMention
if err := rows.Scan(&item.Kind, &item.ID, &item.Label, &item.Description); err != nil {
return page, err
}
page.Items = append(page.Items, item)
}
if err := rows.Err(); err != nil {
return page, err
}
if len(page.Items) > 20 {
page.Items = page.Items[:20]
last := page.Items[19]
blob, _ := json.Marshal(chatMentionCursor{last.ID, last.Kind, fmt.Sprint(last.ID) == query, kind, query})
page.NextCursor = base64.RawURLEncoding.EncodeToString(blob)
}
return page, nil
}
+353
View File
@@ -0,0 +1,353 @@
package db
import (
"database/sql"
"fmt"
"time"
)
// CommandRecord is a paired tool_use + tool_result from the activity table
// (any tool, not just Bash). Command holds the raw tool input (JSON).
type CommandRecord struct {
ID int64 `json:"id"`
ExpID int64 `json:"exploration_id"`
Worker string `json:"worker"`
Tool string `json:"tool"`
Command string `json:"command"`
Output string `json:"output"`
IsError bool `json:"is_error"`
CreatedAt time.Time `json:"created_at"`
}
// commandFilter builds the WHERE clause shared by the tool-execution list and
// its per-tool tally, so the summary always describes exactly the rows the table
// pages through. Returns the clause, its args, and the next placeholder index.
func commandFilter(expID *int64, q string) (string, []any, int) {
where := `WHERE u.kind = 'tool_use'`
args := []any{}
argN := 1
if expID != nil {
where += fmt.Sprintf(` AND u.exploration_id = $%d`, argN)
args = append(args, *expID)
argN++
}
if q != "" {
where += fmt.Sprintf(` AND (u.tool ILIKE $%d OR u.detail ILIKE $%d)`, argN, argN)
args = append(args, "%"+q+"%")
argN++
}
return where, args, argN
}
// ToolStat is one tool's execution tally for the usage summary.
type ToolStat struct {
Tool string `json:"tool"`
Total int `json:"total"`
Errors int `json:"errors"`
}
// ToolStats counts executions grouped by tool under the same filters
// ListCommands takes. Unpaginated on purpose: the tally describes the whole
// filtered set, not the page currently on screen.
func (d *DB) ToolStats(expID *int64, q string) ([]ToolStat, error) {
where, args, _ := commandFilter(expID, q)
rows, err := d.Query(`
SELECT COALESCE(NULLIF(u.tool,''),'-') AS tool, COUNT(*) AS total,
COUNT(*) FILTER (WHERE COALESCE(r.is_error,false)) AS errors
FROM activity u
LEFT JOIN activity r ON r.tool_use_id = u.tool_use_id AND r.kind = 'tool_result'
`+where+`
GROUP BY 1
ORDER BY total DESC, tool ASC`, args...)
if err != nil {
return nil, err
}
defer rows.Close()
out := []ToolStat{}
for rows.Next() {
var s ToolStat
if err := rows.Scan(&s.Tool, &s.Total, &s.Errors); err != nil {
return nil, err
}
out = append(out, s)
}
return out, rows.Err()
}
// ListCommands returns tool executions (tool_use + paired tool_result) across all
// explorations, with optional filtering and pagination. Covers every tool, not
// just Bash; q matches the tool name or its input.
func (d *DB) ListCommands(expID *int64, q string, page, size int) ([]CommandRecord, int, error) {
if size <= 0 {
size = 50
}
if page < 0 {
page = 0
}
offset := page * size
where, args, argN := commandFilter(expID, q)
// count
var total int
countQ := `SELECT COUNT(*) FROM activity u ` + where
if err := d.QueryRow(countQ, args...).Scan(&total); err != nil {
return nil, 0, err
}
// data query: join tool_use with its tool_result
dataQ := `
SELECT u.id, u.exploration_id, COALESCE(u.worker,''), COALESCE(u.tool,''), COALESCE(u.detail,''),
COALESCE(r.detail,''), COALESCE(r.is_error, false), u.created_at
FROM activity u
LEFT JOIN activity r ON r.tool_use_id = u.tool_use_id AND r.kind = 'tool_result'
` + where + `
ORDER BY u.id DESC
LIMIT $` + fmt.Sprintf("%d", argN) + ` OFFSET $` + fmt.Sprintf("%d", argN+1)
args = append(args, size, offset)
rows, err := d.Query(dataQ, args...)
if err != nil {
return nil, 0, err
}
defer rows.Close()
out := []CommandRecord{}
for rows.Next() {
var c CommandRecord
if err := rows.Scan(&c.ID, &c.ExpID, &c.Worker, &c.Tool, &c.Command, &c.Output, &c.IsError, &c.CreatedAt); err != nil {
return nil, 0, err
}
out = append(out, c)
}
return out, total, rows.Err()
}
// LLMRecord is one recorded LLM API call (request + response).
type LLMRecord struct {
ID int64 `json:"id"`
Ts time.Time `json:"ts"`
Model string `json:"model"`
ProfileName string `json:"profile_name"`
SessionID string `json:"session_id"`
TaskID string `json:"task_id"`
Worker string `json:"worker"`
LatencyMs int `json:"latency_ms"`
InputTokens int `json:"input_tokens"`
OutputTokens int `json:"output_tokens"`
CacheRead int `json:"cache_read"`
CacheWrite int `json:"cache_write"`
Status string `json:"status"`
Error string `json:"error,omitempty"`
RequestBody string `json:"request_body,omitempty"`
ResponseBody string `json:"response_body,omitempty"`
// RawRequest / RawResponse are the untouched HTTP bodies exchanged with the
// provider — the request as buildBody() sent it (full tool schemas included)
// and the raw SSE frames. RequestBody/ResponseBody above are the normalized
// view, which drops tool schemas and tool_use blocks entirely. Empty for
// records written before this was added, or when the call never reached HTTP.
RawRequest string `json:"raw_request,omitempty"`
RawResponse string `json:"raw_response,omitempty"`
}
const llmRecordsSchema = `
CREATE TABLE IF NOT EXISTS llm_records (
id BIGSERIAL PRIMARY KEY,
ts TIMESTAMPTZ DEFAULT now(),
model TEXT,
profile_name TEXT,
session_id TEXT,
task_id TEXT,
worker TEXT,
latency_ms INTEGER,
input_tokens INTEGER,
output_tokens INTEGER,
cache_read INTEGER,
cache_write INTEGER,
status TEXT,
error TEXT,
request_body TEXT,
response_body TEXT,
raw_request TEXT,
raw_response TEXT
);
CREATE INDEX IF NOT EXISTS idx_llm_records_ts ON llm_records(ts);
CREATE INDEX IF NOT EXISTS idx_llm_records_session ON llm_records(session_id);
`
// llmRecordsMigrate adds new columns to existing tables.
const llmRecordsMigrate = `
ALTER TABLE llm_records ADD COLUMN IF NOT EXISTS task_id TEXT;
ALTER TABLE llm_records ADD COLUMN IF NOT EXISTS worker TEXT;
ALTER TABLE llm_records ADD COLUMN IF NOT EXISTS profile_name TEXT;
ALTER TABLE llm_records ADD COLUMN IF NOT EXISTS raw_request TEXT;
ALTER TABLE llm_records ADD COLUMN IF NOT EXISTS raw_response TEXT;
`
// EnsureLLMRecordsTable creates the llm_records table if it does not exist.
func (d *DB) EnsureLLMRecordsTable() error {
tx, err := d.Begin()
if err != nil {
return err
}
defer tx.Rollback() //nolint:errcheck
if err := coordinateWithSchemaMigration(tx); err != nil {
return err
}
if _, err := tx.Exec(llmRecordsSchema); err != nil {
return err
}
if _, err := tx.Exec(llmRecordsMigrate); err != nil {
return err
}
return tx.Commit()
}
// InsertLLMRecord stores one LLM call record.
func (d *DB) InsertLLMRecord(r *LLMRecord) error {
_, err := d.Exec(`
INSERT INTO llm_records(model, profile_name, session_id, task_id, worker, latency_ms, input_tokens, output_tokens, cache_read, cache_write, status, error, request_body, response_body, raw_request, raw_response)
VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16)`,
r.Model, nullIfEmpty(r.ProfileName), r.SessionID, nullIfEmpty(r.TaskID), nullIfEmpty(r.Worker),
r.LatencyMs, r.InputTokens, r.OutputTokens, r.CacheRead, r.CacheWrite,
r.Status, nullIfEmpty(r.Error), nullIfEmpty(r.RequestBody), nullIfEmpty(r.ResponseBody),
nullIfEmpty(r.RawRequest), nullIfEmpty(r.RawResponse))
return err
}
// ListLLMRecords returns paginated LLM records with optional filters.
func (d *DB) ListLLMRecords(model, session, task string, page, size int) ([]LLMRecord, int, error) {
if size <= 0 {
size = 50
}
if page < 0 {
page = 0
}
offset := page * size
where := `WHERE true`
args := []any{}
argN := 1
if model != "" {
where += fmt.Sprintf(` AND model = $%d`, argN)
args = append(args, model)
argN++
}
if session != "" {
where += fmt.Sprintf(` AND session_id ILIKE $%d`, argN)
args = append(args, "%"+session+"%")
argN++
}
if task != "" {
where += fmt.Sprintf(` AND COALESCE(task_id,'') = $%d`, argN)
args = append(args, task)
argN++
}
var total int
if err := d.QueryRow(`SELECT COUNT(*) FROM llm_records `+where, args...).Scan(&total); err != nil {
return nil, 0, err
}
dataQ := `SELECT id, ts, COALESCE(model,''), COALESCE(profile_name,''), COALESCE(session_id,''), COALESCE(task_id,''), COALESCE(worker,''),
COALESCE(latency_ms,0), COALESCE(input_tokens,0), COALESCE(output_tokens,0), COALESCE(cache_read,0), COALESCE(cache_write,0),
COALESCE(status,''), COALESCE(error,'')
FROM llm_records ` + where + ` ORDER BY id DESC LIMIT $` + fmt.Sprintf("%d", argN) + ` OFFSET $` + fmt.Sprintf("%d", argN+1)
args = append(args, size, offset)
rows, err := d.Query(dataQ, args...)
if err != nil {
return nil, 0, err
}
defer rows.Close()
out := []LLMRecord{}
for rows.Next() {
var r LLMRecord
if err := rows.Scan(&r.ID, &r.Ts, &r.Model, &r.ProfileName, &r.SessionID, &r.TaskID, &r.Worker, &r.LatencyMs,
&r.InputTokens, &r.OutputTokens, &r.CacheRead, &r.CacheWrite, &r.Status, &r.Error); err != nil {
return nil, 0, err
}
out = append(out, r)
}
return out, total, rows.Err()
}
// ModelTokenStat is one model's aggregated token usage for a task, summed from the
// llm_usage metering ledger (see db/llm_usage.go). Calls is the number of LLM calls
// that hit this model.
type ModelTokenStat struct {
Model string `json:"model"`
Calls int `json:"calls"`
InputTokens int `json:"input_tokens"`
OutputTokens int `json:"output_tokens"`
CacheReadTokens int `json:"cache_read_tokens"`
CacheWriteTokens int `json:"cache_write_tokens"`
}
// LLMTask is one distinct task with its LLM-record count.
type LLMTask struct {
TaskID string `json:"task_id"`
Count int `json:"count"`
}
// LLMTasks returns distinct non-empty task_ids with record counts, most recent
// first — powers the LLM-records page's task picker.
func (d *DB) LLMTasks() ([]LLMTask, error) {
rows, err := d.Query(`SELECT task_id, COUNT(*) AS n FROM llm_records
WHERE COALESCE(task_id,'') <> '' GROUP BY task_id ORDER BY MAX(id) DESC`)
if err != nil {
return nil, err
}
defer rows.Close()
var out []LLMTask
for rows.Next() {
var t LLMTask
if err := rows.Scan(&t.TaskID, &t.Count); err != nil {
return nil, err
}
out = append(out, t)
}
return out, rows.Err()
}
// DeleteLLMRecords removes every LLM record for one exact task_id — the same
// match the page's task picker/filter uses. Returns rows deleted.
func (d *DB) DeleteLLMRecords(task string) (int64, error) {
res, err := d.Exec(`DELETE FROM llm_records WHERE COALESCE(task_id,'') = $1`, task)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
// GetLLMRecord returns a single LLM record with full request/response bodies.
func (d *DB) GetLLMRecord(id int64) (*LLMRecord, error) {
var r LLMRecord
var reqBody, respBody, rawReq, rawResp sql.NullString
err := d.QueryRow(`SELECT id, ts, COALESCE(model,''), COALESCE(profile_name,''), COALESCE(session_id,''), COALESCE(task_id,''), COALESCE(worker,''),
COALESCE(latency_ms,0), COALESCE(input_tokens,0), COALESCE(output_tokens,0), COALESCE(cache_read,0), COALESCE(cache_write,0),
COALESCE(status,''), COALESCE(error,''), request_body, response_body, raw_request, raw_response
FROM llm_records WHERE id=$1`, id).
Scan(&r.ID, &r.Ts, &r.Model, &r.ProfileName, &r.SessionID, &r.TaskID, &r.Worker, &r.LatencyMs,
&r.InputTokens, &r.OutputTokens, &r.CacheRead, &r.CacheWrite, &r.Status, &r.Error,
&reqBody, &respBody, &rawReq, &rawResp)
if err != nil {
return nil, err
}
r.RequestBody = reqBody.String
r.ResponseBody = respBody.String
r.RawRequest = rawReq.String
r.RawResponse = rawResp.String
return &r, nil
}
func nullIfEmpty(s string) any {
if s == "" {
return nil
}
return s
}
+774
View File
@@ -0,0 +1,774 @@
package db
import (
"database/sql"
"errors"
"fmt"
"log"
"net"
"strings"
"unicode/utf8"
)
// =====================================================================
// 公司主体层
// =====================================================================
// Company is a row in the companies table.
type Company struct {
ID int64 `json:"id"`
Name string `json:"name"`
NKey string `json:"nkey"`
Logo *string `json:"logo,omitempty"`
CreatedAt string `json:"created_at"`
UpdatedAt string `json:"updated_at"`
}
// CompanyWithScope extends Company with its scope rules and asset count.
type CompanyWithScope struct {
Company
Scope []ScopeRule `json:"scope"`
AssetCount int `json:"asset_count"`
}
// ScopeRule is one company_scope row.
type ScopeRule struct {
ID int64 `json:"id"`
CompanyID int64 `json:"company_id"`
Kind string `json:"kind"`
Domain string `json:"domain,omitempty"`
Net string `json:"net,omitempty"`
Value string `json:"value,omitempty"`
Raw string `json:"raw"`
Reason string `json:"reason,omitempty"`
}
// CompanyStore operates on the companies + company_scope tables.
type CompanyStore struct{ db *DB }
var (
ErrCompanyNameConflict = errors.New("company name already exists")
ErrCompanyNotFound = errors.New("company not found")
)
const (
// 企业范围不限制规则条数:逐个 IP / 域名录入的范围动辄上千条,封顶只会逼用户
// 拆成多个企业。请求体大小(server 侧 maxCompanyMutationBodyBytes)仍然兜底。
//
// Raw and normalized textual scope payloads are bounded by Unicode rune
// count so multi-byte input is treated consistently by the API and DB layer.
MaxCompanyScopeRawRunes = 1024
MaxCompanyScopeValueRunes = 1024
)
// CompanyScopeValidationError identifies a client-correctable scope error.
// Storage and transaction failures are returned as ordinary errors instead.
type CompanyScopeValidationError struct{ Message string }
func (e *CompanyScopeValidationError) Error() string { return e.Message }
// ValidateCompanyScopeInputBounds applies request-wide limits before parsing.
// Store methods call it again so non-HTTP callers cannot bypass the limits.
// 只约束单条规则的长度,不限制条数。
func ValidateCompanyScopeInputBounds(inputs []ScopeInput) error {
for i, input := range inputs {
if utf8.RuneCountInString(input.Value) > MaxCompanyScopeRawRunes {
return &CompanyScopeValidationError{Message: fmt.Sprintf(
"企业范围第 %d 条原始值过长: 最多 %d 个字符", i+1, MaxCompanyScopeRawRunes,
)}
}
}
return nil
}
// Scope writes rebuild derived asset ownership globally, so serialize them to
// ensure the committed attribution always reflects the latest committed rules.
// This key is reserved for company mutations; 7337741001 is the schema lock and
// 7337741002 is the cross-package test-suite lock.
const companyScopeMutationLock int64 = 7337741003
// Companies returns the company store.
func (d *DB) Companies() *CompanyStore { return &CompanyStore{db: d} }
// companyNKey normalises a company name: lowercase + trim + collapse whitespace.
func companyNKey(name string) string {
return strings.Join(strings.Fields(strings.ToLower(name)), " ")
}
// UpsertCompany creates or updates a company by name. Returns the id and whether
// a new row was created.
func (s *CompanyStore) UpsertCompany(name, logo string) (id int64, created bool, err error) {
nkey := companyNKey(name)
var logoVal any
if logo != "" {
logoVal = logo
}
err = s.db.QueryRow(`
INSERT INTO companies(name, nkey, logo)
VALUES ($1, $2, $3)
ON CONFLICT (nkey) DO UPDATE SET
name = EXCLUDED.name,
logo = COALESCE(EXCLUDED.logo, companies.logo),
updated_at = now()
RETURNING id, (xmax = 0)`, name, nkey, logoVal).Scan(&id, &created)
return
}
// CreateCompanyWithScope creates a company without updating an existing row.
// The company, its valid initial scope rules, and derived asset attribution are
// committed atomically. Invalid inputs retain the legacy partial-validation
// contract and are reported without preventing valid rules from being stored.
func (s *CompanyStore) CreateCompanyWithScope(name, logo string, inputs []ScopeInput, reason string) (
id int64, added, skipped, invalid int, validationErrors []string, err error,
) {
if err := ValidateCompanyScopeInputBounds(inputs); err != nil {
return 0, 0, 0, 0, nil, err
}
rules, invalid, validationErrors := parseScopeInputs(inputs)
if err := validateParsedScopeBounds(rules); err != nil {
return 0, 0, 0, invalid, validationErrors, err
}
tx, err := s.db.Begin()
if err != nil {
return 0, 0, 0, invalid, validationErrors, err
}
defer tx.Rollback() //nolint:errcheck
if err := lockCompanyScopeMutation(tx); err != nil {
return 0, 0, 0, invalid, validationErrors, err
}
nkey := companyNKey(name)
var logoVal any
if logo != "" {
logoVal = logo
}
if err := tx.QueryRow(`
INSERT INTO companies(name, nkey, logo)
VALUES ($1, $2, $3)
ON CONFLICT (nkey) DO NOTHING
RETURNING id`, name, nkey, logoVal).Scan(&id); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return 0, 0, 0, invalid, validationErrors, ErrCompanyNameConflict
}
return 0, 0, 0, invalid, validationErrors, err
}
added, skipped, needsAttribution, err := insertScopeRulesTx(tx, id, rules, reason)
if err != nil {
return 0, 0, 0, invalid, validationErrors, err
}
if needsAttribution {
warning, err := recomputeAttributionTx(tx)
if err != nil {
return 0, 0, 0, invalid, validationErrors, err
}
logAttributionWarning(warning)
}
if err := tx.Commit(); err != nil {
return 0, 0, 0, invalid, validationErrors, err
}
return id, added, skipped, invalid, validationErrors, nil
}
// GetCompany returns one company by id (nil if not found).
func (s *CompanyStore) GetCompany(id int64) (*Company, error) {
c := &Company{}
err := s.db.QueryRow(`
SELECT id, name, nkey, logo, created_at::text, updated_at::text
FROM companies WHERE id = $1`, id).Scan(
&c.ID, &c.Name, &c.NKey, &c.Logo, &c.CreatedAt, &c.UpdatedAt)
if err == sql.ErrNoRows {
return nil, nil
}
return c, err
}
// GetCompanyByName returns one company by normalized name (nil if not found).
func (s *CompanyStore) GetCompanyByName(name string) (*Company, error) {
nkey := companyNKey(name)
c := &Company{}
err := s.db.QueryRow(`
SELECT id, name, nkey, logo, created_at::text, updated_at::text
FROM companies WHERE nkey = $1`, nkey).Scan(
&c.ID, &c.Name, &c.NKey, &c.Logo, &c.CreatedAt, &c.UpdatedAt)
if err == sql.ErrNoRows {
return nil, nil
}
return c, err
}
// UpsertByName creates the company if it doesn't exist, then returns its id.
func (s *CompanyStore) UpsertByName(name string) (int64, error) {
id, _, err := s.UpsertCompany(name, "")
return id, err
}
// DeleteCompany deletes a company and re-evaluates automatic ownership against
// the remaining companies in the same transaction. Explicitly-owned assets are
// detached by the FK and may then fall back to a remaining scope match.
func (s *CompanyStore) DeleteCompany(id int64) error {
_, err := s.DeleteCompanyWithAssets(id, false)
return err
}
// DeleteCompanyWithAssets deletes a company and optionally all of its assets in
// one transaction, then re-evaluates ownership against the remaining companies.
func (s *CompanyStore) DeleteCompanyWithAssets(id int64, deleteAssets bool) (assetsDeleted int64, err error) {
tx, err := s.db.Begin()
if err != nil {
return 0, err
}
defer tx.Rollback() //nolint:errcheck
if err := lockCompanyScopeMutation(tx); err != nil {
return 0, err
}
if deleteAssets {
res, err := tx.Exec(`DELETE FROM assets WHERE company_id = $1`, id)
if err != nil {
return 0, err
}
assetsDeleted, err = res.RowsAffected()
if err != nil {
return 0, err
}
}
res, err := tx.Exec(`DELETE FROM companies WHERE id = $1`, id)
if err != nil {
return 0, err
}
deleted, err := res.RowsAffected()
if err != nil {
return 0, err
}
if deleted == 0 {
return 0, ErrCompanyNotFound
}
// This path has no per-request warning channel, so the log is the only place
// the operator can learn about unparseable ip rows here.
warning, err := recomputeAttributionTx(tx)
if err != nil {
return 0, err
}
logAttributionWarning(warning)
if err := tx.Commit(); err != nil {
return 0, err
}
return assetsDeleted, nil
}
// ListCompanies returns all companies with scope and asset count.
func (s *CompanyStore) ListCompanies() ([]*CompanyWithScope, error) {
rows, err := s.db.Query(`
SELECT c.id, c.name, c.nkey, c.logo, c.created_at::text, c.updated_at::text,
COUNT(DISTINCT a.id) AS asset_count
FROM companies c
LEFT JOIN assets a ON a.company_id = c.id
GROUP BY c.id
ORDER BY c.name`)
if err != nil {
return nil, err
}
defer rows.Close()
var out []*CompanyWithScope
for rows.Next() {
cws := &CompanyWithScope{}
if err := rows.Scan(&cws.ID, &cws.Name, &cws.NKey, &cws.Logo,
&cws.CreatedAt, &cws.UpdatedAt, &cws.AssetCount); err != nil {
return nil, err
}
out = append(out, cws)
}
if err := rows.Err(); err != nil {
return nil, err
}
// fetch scope rules for each company
for _, cws := range out {
cws.Scope, err = s.GetScope(cws.ID)
if err != nil {
return nil, err
}
}
return out, nil
}
// GetScope returns all scope rules for a company.
func (s *CompanyStore) GetScope(companyID int64) ([]ScopeRule, error) {
rows, err := s.db.Query(`
SELECT id, company_id, kind,
COALESCE(domain,''), COALESCE(net::text,''), COALESCE(value,''), raw, COALESCE(reason,'')
FROM company_scope
WHERE company_id = $1
ORDER BY id`, companyID)
if err != nil {
return nil, err
}
defer rows.Close()
out := make([]ScopeRule, 0)
for rows.Next() {
var r ScopeRule
if err := rows.Scan(&r.ID, &r.CompanyID, &r.Kind, &r.Domain, &r.Net, &r.Value, &r.Raw, &r.Reason); err != nil {
return nil, err
}
out = append(out, r)
}
return out, rows.Err()
}
// AddScope parses and inserts scope lines for a company, then reattributes assets.
// Returns counts of added, skipped, and invalid lines.
func (s *CompanyStore) AddScope(companyID int64, lines []string, reason string) (added, skipped, invalid int, errors []string) {
inputs := make([]ScopeInput, 0, len(lines))
for _, line := range lines {
inputs = append(inputs, ScopeInput{Value: line})
}
return s.AddScopeInputs(companyID, inputs, reason)
}
// AddScopeInputs inserts structured scope rules. Empty kinds use the automatic
// CIDR/IP/ICP/domain/keyword classification used by AddScope.
func (s *CompanyStore) AddScopeInputs(companyID int64, inputs []ScopeInput, reason string) (added, skipped, invalid int, errors []string) {
added, skipped, invalid, validationErrors, err := s.AddScopeInputsChecked(companyID, inputs, reason)
if err != nil {
validationErrors = append(validationErrors, err.Error())
}
return added, skipped, invalid, validationErrors
}
// AddScopeInputsChecked inserts structured scope rules while keeping input
// validation separate from storage and transaction errors.
func (s *CompanyStore) AddScopeInputsChecked(companyID int64, inputs []ScopeInput, reason string) (
added, skipped, invalid int, validationErrors []string, err error,
) {
if err := ValidateCompanyScopeInputBounds(inputs); err != nil {
return 0, 0, 0, nil, err
}
rules, invalid, errors := parseScopeInputs(inputs)
if err := validateParsedScopeBounds(rules); err != nil {
return 0, 0, invalid, errors, err
}
tx, err := s.db.Begin()
if err != nil {
return 0, 0, invalid, errors, err
}
defer tx.Rollback() //nolint:errcheck
if err := lockCompanyScopeMutation(tx); err != nil {
return 0, 0, invalid, errors, err
}
if err := ensureCompanyExistsTx(tx, companyID); err != nil {
return 0, 0, invalid, errors, err
}
if len(rules) == 0 {
if err := tx.Commit(); err != nil {
return 0, 0, invalid, errors, err
}
return 0, 0, invalid, errors, nil
}
added, skipped, needsAttribution, err := insertScopeRulesTx(tx, companyID, rules, reason)
if err != nil {
return 0, 0, invalid, errors, err
}
if needsAttribution {
warning, err := recomputeAttributionTx(tx)
if err != nil {
return 0, 0, invalid, errors, fmt.Errorf("重新计算企业归属失败: %w", err)
}
logAttributionWarning(warning)
}
if err := tx.Commit(); err != nil {
return 0, 0, invalid, errors, err
}
return added, skipped, invalid, errors, nil
}
func parseScopeInputs(inputs []ScopeInput) (rules []ParsedScope, invalid int, validationErrors []string) {
rules = make([]ParsedScope, 0, len(inputs))
for _, input := range inputs {
rule, err := ParseScopeInput(input)
if err != nil {
invalid++
validationErrors = append(validationErrors, fmt.Sprintf("%s: %v", input.Value, err))
continue
}
rules = append(rules, rule)
}
return rules, invalid, validationErrors
}
func validateParsedScopeBounds(rules []ParsedScope) error {
for i, rule := range rules {
if utf8.RuneCountInString(rule.Raw) > MaxCompanyScopeRawRunes {
return &CompanyScopeValidationError{Message: fmt.Sprintf(
"企业范围第 %d 条原始值过长: 最多 %d 个字符", i+1, MaxCompanyScopeRawRunes,
)}
}
if utf8.RuneCountInString(rule.Value) > MaxCompanyScopeValueRunes {
return &CompanyScopeValidationError{Message: fmt.Sprintf(
"企业范围第 %d 条规范化值过长: 最多 %d 个字符", i+1, MaxCompanyScopeValueRunes,
)}
}
}
return nil
}
func ensureCompanyExistsTx(tx *sql.Tx, companyID int64) error {
var exists bool
if err := tx.QueryRow(`SELECT EXISTS(SELECT 1 FROM companies WHERE id = $1)`, companyID).Scan(&exists); err != nil {
return err
}
if !exists {
return ErrCompanyNotFound
}
return nil
}
func lockCompanyScopeMutation(tx *sql.Tx) error {
_, err := tx.Exec(`SELECT pg_advisory_xact_lock($1)`, companyScopeMutationLock)
return err
}
func insertScopeRulesTx(tx *sql.Tx, companyID int64, rules []ParsedScope, reason string) (
added, skipped int, needsAttribution bool, err error,
) {
for _, rule := range rules {
inserted, insertErr := insertScopeRuleTx(tx, companyID, rule, reason)
if insertErr != nil {
return 0, 0, false, insertErr
}
if !inserted {
skipped++
continue
}
added++
needsAttribution = needsAttribution || rule.Kind != "keyword"
}
return added, skipped, needsAttribution, nil
}
// insertScopeRuleTx inserts a scope rule. inserted=false means a duplicate was
// ignored by ON CONFLICT, not an error.
func insertScopeRuleTx(tx *sql.Tx, companyID int64, rule ParsedScope, reason string) (inserted bool, err error) {
var res interface{ RowsAffected() (int64, error) }
switch rule.Kind {
case "domain":
res, err = tx.Exec(`
INSERT INTO company_scope(company_id, kind, domain, raw, reason)
VALUES ($1, 'domain', $2, $3, $4)
ON CONFLICT ON CONSTRAINT uq_sv2_domain DO NOTHING`,
companyID, rule.Domain, rule.Raw, reason)
case "ip", "cidr":
res, err = tx.Exec(`
INSERT INTO company_scope(company_id, kind, net, raw, reason)
VALUES ($1, $2, $3::cidr, $4, $5)
ON CONFLICT ON CONSTRAINT uq_sv2_net DO NOTHING`,
companyID, rule.Kind, rule.Net, rule.Raw, reason)
case "icp", "keyword":
res, err = tx.Exec(`
INSERT INTO company_scope(company_id, kind, value, raw, reason)
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (company_id, kind, value) WHERE kind IN ('icp','keyword') DO NOTHING`,
companyID, rule.Kind, rule.Value, rule.Raw, reason)
default:
return false, fmt.Errorf("unsupported company scope kind %q", rule.Kind)
}
if err != nil {
return false, err
}
n, _ := res.RowsAffected()
return n > 0, nil
}
// RecomputeAttribution rebuilds only scope-derived ownership. Explicit company
// links are immutable under scope edits. Precedence is domain, IP/CIDR, then
// normalized exact ICP; keyword rules never attribute assets.
func (s *CompanyStore) RecomputeAttribution() error {
tx, err := s.db.Begin()
if err != nil {
return err
}
defer tx.Rollback() //nolint:errcheck
if err := lockCompanyScopeMutation(tx); err != nil {
return err
}
warning, err := recomputeAttributionTx(tx)
if err != nil {
return err
}
logAttributionWarning(warning)
return tx.Commit()
}
// recomputeAttributionTx rebuilds scope-derived ownership. It returns a warning
// for assets whose ip column cannot be parsed: try_inet skips them instead of
// aborting the statement, so without this they would silently never receive a
// network-based company. Callers surface the warning and it is always logged.
func recomputeAttributionTx(tx *sql.Tx) (string, error) {
// Only derived rows are cleared. Historical rows migrated without provenance
// are marked explicit by schema.sql, which is the non-destructive default.
if _, err := tx.Exec(`
UPDATE assets
SET company_id = NULL, company_source = 'scope'
WHERE company_source = 'scope'`); err != nil {
return "", err
}
// Domain-based attribution (root_domain exact match).
if _, err := tx.Exec(`
WITH matched AS (
SELECT DISTINCT ON (a.id) a.id AS asset_id, cs.company_id
FROM assets a
JOIN company_scope cs ON cs.kind = 'domain' AND a.root_domain = cs.domain
WHERE a.company_id IS NULL
AND a.type IN ('root_domain','subdomain','service','endpoint')
AND a.root_domain IS NOT NULL
ORDER BY a.id, length(cs.domain) DESC, cs.company_id
)
UPDATE assets a
SET company_id = matched.company_id, company_source = 'scope'
FROM matched
WHERE a.id = matched.asset_id`); err != nil {
return "", err
}
// IP/CIDR attribution for still-unowned assets.
if _, err := tx.Exec(`
WITH matched AS (
SELECT DISTINCT ON (a.id) a.id AS asset_id, cs.company_id
FROM assets a
JOIN company_scope cs ON cs.kind IN ('ip','cidr') AND cs.net >>= try_inet(a.ip)
WHERE a.company_id IS NULL
AND a.type IN ('ip','subdomain','service','endpoint')
AND a.ip IS NOT NULL
ORDER BY a.id, masklen(cs.net) DESC, cs.company_id
)
UPDATE assets a
SET company_id = matched.company_id, company_source = 'scope'
FROM matched
WHERE a.id = matched.asset_id`); err != nil {
return "", err
}
// Exact normalized ICP attribution after domain/network precedence.
if _, err := tx.Exec(`
WITH matched AS (
SELECT DISTINCT ON (a.id) a.id AS asset_id, cs.company_id
FROM assets a
JOIN company_scope cs ON cs.kind = 'icp'
AND (
lower(regexp_replace(COALESCE(a.icp,''), '[[:space:]]+', '', 'g')) = cs.value
OR lower(regexp_replace(COALESCE(a.app_icp,''), '[[:space:]]+', '', 'g')) = cs.value
)
WHERE a.company_id IS NULL
AND (COALESCE(a.icp,'') <> '' OR COALESCE(a.app_icp,'') <> '')
ORDER BY a.id, cs.company_id
)
UPDATE assets a
SET company_id = matched.company_id, company_source = 'scope'
FROM matched
WHERE a.id = matched.asset_id`); err != nil {
return "", err
}
return malformedIPAssetWarning(tx)
}
// malformedIPAssetsSampled bounds how many offending ids one warning names, so a
// large batch of bad rows stays readable in a toast and in the log.
const malformedIPAssetsSampled = 5
// logAttributionWarning records a recompute warning in the server log. Every
// recompute path calls it, so the warning is reported even for triggers with no
// per-request response (company deletion, scopesentry sync, agent asset writes).
func logAttributionWarning(warning string) {
if warning != "" {
log.Printf("[assets] %s", warning)
}
}
// malformedIPAssetQueryer is satisfied by both *sql.Tx and *DB so the warning
// can be produced inside a recompute transaction or read standalone by the API.
type malformedIPAssetQueryer interface {
Query(query string, args ...any) (*sql.Rows, error)
}
// MalformedIPAssetWarning reports assets with an unparseable ip outside of any
// mutation, letting the API attach the warning to a scope response without
// widening the mutation signatures — an unrelated data problem is not one of
// this request's validation errors.
func (s *CompanyStore) MalformedIPAssetWarning() (string, error) {
return malformedIPAssetWarning(s.db)
}
// malformedIPAssetWarning describes assets whose ip column is not a valid
// address. They are invisible to network attribution, so the operator has to be
// told which rows to fix — silently skipping them would look like scope rules
// that simply do not work.
func malformedIPAssetWarning(q malformedIPAssetQueryer) (string, error) {
rows, err := q.Query(`
SELECT id, ip, count(*) OVER () AS total
FROM assets
WHERE ip IS NOT NULL AND ip <> '' AND try_inet(ip) IS NULL
AND type IN ('ip','subdomain','service','endpoint')
ORDER BY id
LIMIT $1`, malformedIPAssetsSampled)
if err != nil {
return "", err
}
defer rows.Close()
var total int
samples := make([]string, 0, malformedIPAssetsSampled)
for rows.Next() {
var id int64
var ip string
if err := rows.Scan(&id, &ip, &total); err != nil {
return "", err
}
samples = append(samples, fmt.Sprintf("#%d %s", id, ip))
}
if err := rows.Err(); err != nil {
return "", err
}
if total == 0 {
return "", nil
}
warning := fmt.Sprintf(
"%d 条资产的 ip 字段不是合法 IP,已跳过 IP/CIDR 范围匹配(这些资产不会被网段规则归属到企业):%s",
total, strings.Join(samples, "、"),
)
if total > len(samples) {
warning += fmt.Sprintf(" 等 %d 条", total)
}
return warning, nil
}
// UpdateScope replaces all scope rules for a company and reattributes.
func (s *CompanyStore) UpdateScope(companyID int64, lines []string, reason string) (added, invalid int, errs []string) {
inputs := make([]ScopeInput, 0, len(lines))
for _, line := range lines {
inputs = append(inputs, ScopeInput{Value: line})
}
return s.UpdateScopeInputs(companyID, inputs, reason)
}
// UpdateScopeInputs replaces all rules with a structured set.
func (s *CompanyStore) UpdateScopeInputs(companyID int64, inputs []ScopeInput, reason string) (added, invalid int, errs []string) {
added, invalid, validationErrors, err := s.UpdateScopeInputsChecked(companyID, inputs, reason)
if err != nil {
validationErrors = append(validationErrors, err.Error())
}
return added, invalid, validationErrors
}
// UpdateScopeInputsChecked replaces all rules while separating validation
// feedback from storage and transaction failures.
func (s *CompanyStore) UpdateScopeInputsChecked(companyID int64, inputs []ScopeInput, reason string) (
added, invalid int, validationErrors []string, err error,
) {
if err := ValidateCompanyScopeInputBounds(inputs); err != nil {
return 0, 0, nil, err
}
rules, invalid, errs := parseScopeInputs(inputs)
if invalid > 0 {
return 0, invalid, errs, &CompanyScopeValidationError{Message: fmt.Sprintf(
"企业范围包含 %d 条无效规则,未覆盖原有范围", invalid,
)}
}
if err := validateParsedScopeBounds(rules); err != nil {
return 0, invalid, errs, err
}
tx, err := s.db.Begin()
if err != nil {
return 0, invalid, errs, err
}
defer tx.Rollback() //nolint:errcheck
if err := lockCompanyScopeMutation(tx); err != nil {
return 0, invalid, errs, err
}
if err := ensureCompanyExistsTx(tx, companyID); err != nil {
return 0, invalid, errs, err
}
if _, err := tx.Exec(`DELETE FROM company_scope WHERE company_id = $1`, companyID); err != nil {
return 0, invalid, errs, err
}
added, _, _, err = insertScopeRulesTx(tx, companyID, rules, reason)
if err != nil {
return 0, invalid, errs, err
}
// Rebuild even for an empty replacement because removing the old rules may
// detach scope-derived assets or expose a lower-precedence company match.
warning, err := recomputeAttributionTx(tx)
if err != nil {
return 0, invalid, errs, fmt.Errorf("重新计算企业归属失败: %w", err)
}
logAttributionWarning(warning)
if err := tx.Commit(); err != nil {
return 0, invalid, errs, err
}
return added, invalid, errs, nil
}
// ResolveCompany returns the company_id for a given root_domain and/or ip, or nil
// if no scope rule matches. Mirrors the attribution logic used at asset insert time.
func (s *CompanyStore) ResolveCompany(rootDomain, ipStr string) (*int64, error) {
return s.ResolveCompanyWithICP(rootDomain, ipStr, "")
}
// ResolveCompanyWithICP mirrors RecomputeAttribution for insert-time ownership.
// ICP is consulted only after domain and IP/CIDR fail to match.
func (s *CompanyStore) ResolveCompanyWithICP(rootDomain, ipStr, icp string) (*int64, error) {
return resolveCompanyWithICP(s.db, rootDomain, ipStr, icp)
}
type companyScopeQueryer interface {
QueryRow(query string, args ...any) *sql.Row
}
func resolveCompanyWithICP(q companyScopeQueryer, rootDomain, ipStr, icp string) (*int64, error) {
if rootDomain != "" {
var cid int64
err := q.QueryRow(`
SELECT company_id FROM company_scope
WHERE kind = 'domain'
AND domain = $1
ORDER BY length(domain) DESC, company_id
LIMIT 1`, rootDomain).Scan(&cid)
if err == nil {
return &cid, nil
}
if err != sql.ErrNoRows {
return nil, err
}
}
if ipStr != "" {
if net.ParseIP(ipStr) != nil {
var cid int64
err := q.QueryRow(`
SELECT company_id FROM company_scope
WHERE kind IN ('ip','cidr')
AND net >>= $1::inet
ORDER BY masklen(net) DESC, company_id
LIMIT 1`, ipStr).Scan(&cid)
if err == nil {
return &cid, nil
}
if err != sql.ErrNoRows {
return nil, err
}
}
}
if normalized := NormalizeICP(icp); normalized != "" {
var cid int64
err := q.QueryRow(`
SELECT company_id FROM company_scope
WHERE kind = 'icp' AND value = $1
ORDER BY company_id
LIMIT 1`, normalized).Scan(&cid)
if err == nil {
return &cid, nil
}
if err != sql.ErrNoRows {
return nil, err
}
}
return nil, nil
}
+460
View File
@@ -0,0 +1,460 @@
package db
import (
"errors"
"fmt"
"strings"
"testing"
"time"
)
// cleanup helpers to remove test data
func cleanupCompany(d *DB, id int64) {
d.Exec(`DELETE FROM company_scope WHERE company_id = $1`, id)
d.Exec(`DELETE FROM companies WHERE id = $1`, id)
}
func TestCompanyUpsertAndGet(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v)", err)
}
defer d.Close()
cs := d.Companies()
id, created, err := cs.UpsertCompany("Test Corp", "https://example.com/logo.png")
if err != nil {
t.Fatal(err)
}
defer cleanupCompany(d, id)
if !created {
t.Error("first upsert should report created=true")
}
// duplicate: same nkey, should not create new
id2, created2, err := cs.UpsertCompany("Test Corp", "")
if err != nil {
t.Fatal(err)
}
if id2 != id {
t.Errorf("dedup failed: %d != %d", id2, id)
}
if created2 {
t.Error("second upsert should report created=false")
}
c, err := cs.GetCompany(id)
if err != nil || c == nil {
t.Fatalf("GetCompany: %v", err)
}
if c.Name != "Test Corp" {
t.Errorf("name: %q", c.Name)
}
// GetCompanyByName
c2, err := cs.GetCompanyByName("test corp") // normalised
if err != nil || c2 == nil {
t.Fatalf("GetCompanyByName: %v", err)
}
if c2.ID != id {
t.Errorf("GetCompanyByName id mismatch: %d vs %d", c2.ID, id)
}
}
func TestCreateCompanyWithScopeRejectsNormalizedDuplicate(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v)", err)
}
defer d.Close()
cs := d.Companies()
stamp := time.Now().UnixNano()
name := fmt.Sprintf("Strict Company %d", stamp)
id, added, _, _, validationErrors, err := cs.CreateCompanyWithScope(name, "", []ScopeInput{
{Kind: "domain", Value: fmt.Sprintf("strict-%d.example", stamp)},
}, "test")
if err != nil {
t.Fatal(err)
}
defer cleanupCompany(d, id)
if added != 1 || len(validationErrors) != 0 {
t.Fatalf("initial scope: added=%d errors=%v", added, validationErrors)
}
_, _, _, _, _, err = cs.CreateCompanyWithScope(
" "+strings.ToUpper(strings.ReplaceAll(name, " ", " "))+" ",
"",
[]ScopeInput{{Kind: "domain", Value: fmt.Sprintf("replacement-%d.example", stamp)}},
"test",
)
if !errors.Is(err, ErrCompanyNameConflict) {
t.Fatalf("duplicate create error=%v want ErrCompanyNameConflict", err)
}
scope, err := cs.GetScope(id)
if err != nil {
t.Fatal(err)
}
if len(scope) != 1 || scope[0].Domain != fmt.Sprintf("strict-%d.example", stamp) {
t.Fatalf("duplicate create changed existing scope: %+v", scope)
}
}
func TestCreateCompanyWithScopeRollsBackOnScopeWriteFailure(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v)", err)
}
defer d.Close()
cs := d.Companies()
name := fmt.Sprintf("Atomic Create %d", time.Now().UnixNano())
_, _, _, _, _, err = cs.CreateCompanyWithScope(name, "", []ScopeInput{
{Kind: "keyword", Value: "invalid\x00postgres-text"},
}, "test")
if err == nil {
t.Fatal("expected scope database write to fail")
}
company, getErr := cs.GetCompanyByName(name)
if getErr != nil {
t.Fatal(getErr)
}
if company != nil {
defer cleanupCompany(d, company.ID)
t.Fatalf("company row survived failed initial scope transaction: %+v", company)
}
}
func TestCompanyScope(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v)", err)
}
defer d.Close()
cs := d.Companies()
id, _, err := cs.UpsertCompany("ScopeTestCorp", "")
if err != nil {
t.Fatal(err)
}
defer cleanupCompany(d, id)
lines := []string{"example.com", "192.168.1.0/24", "10.0.0.1"}
added, skipped, invalid, errs := cs.AddScope(id, lines, "test")
if added != 3 {
t.Errorf("want 3 added, got %d (errs: %v)", added, errs)
}
if skipped != 0 || invalid != 0 {
t.Errorf("unexpected skipped=%d invalid=%d", skipped, invalid)
}
// Adding again should skip (duplicate)
added2, skipped2, invalid2, _ := cs.AddScope(id, lines, "test")
if added2 != 0 || skipped2 != 3 {
t.Errorf("want 0 added 3 skipped, got %d added %d skipped %d invalid", added2, skipped2, invalid2)
}
scope, err := cs.GetScope(id)
if err != nil {
t.Fatal(err)
}
if len(scope) != 3 {
t.Errorf("want 3 scope rules, got %d", len(scope))
}
}
func TestCompanyScopeInvalid(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v)", err)
}
defer d.Close()
cs := d.Companies()
id, _, err := cs.UpsertCompany("InvalidScopeCorp", "")
if err != nil {
t.Fatal(err)
}
defer cleanupCompany(d, id)
// Explicitly typed TLD-only domains and overly broad CIDRs should be
// rejected. Untyped plain text is intentionally classified as a keyword.
inputs := []ScopeInput{
{Kind: "domain", Value: "com"},
{Kind: "cidr", Value: "1.2.3.4/8"},
}
added, _, invalid, _ := cs.AddScopeInputs(id, inputs, "test")
if added != 0 {
t.Errorf("want 0 added for invalid lines, got %d", added)
}
if invalid != 2 {
t.Errorf("want 2 invalid, got %d", invalid)
}
}
func TestResolveCompany(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v)", err)
}
defer d.Close()
cs := d.Companies()
id, _, err := cs.UpsertCompany("ResolveCorp", "")
if err != nil {
t.Fatal(err)
}
defer cleanupCompany(d, id)
cs.AddScope(id, []string{"resolve-test.io", "10.20.0.0/16"}, "test")
// domain match
cid, err := cs.ResolveCompany("resolve-test.io", "")
if err != nil || cid == nil || *cid != id {
t.Errorf("domain resolve: want %d, got %v (err %v)", id, cid, err)
}
// no match
cid2, err := cs.ResolveCompany("notinscope.com", "")
if err != nil || cid2 != nil {
t.Errorf("no-match: want nil, got %v", cid2)
}
// IP/CIDR match
cid3, err := cs.ResolveCompany("", "10.20.5.1")
if err != nil || cid3 == nil || *cid3 != id {
t.Errorf("cidr resolve: want %d, got %v (err %v)", id, cid3, err)
}
// IP outside CIDR
cid4, err := cs.ResolveCompany("", "10.30.0.1")
if err != nil || cid4 != nil {
t.Errorf("cidr no-match: want nil, got %v", cid4)
}
}
func TestUpdateScope(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v)", err)
}
defer d.Close()
cs := d.Companies()
id, _, err := cs.UpsertCompany("UpdateScopeCorp", "")
if err != nil {
t.Fatal(err)
}
defer cleanupCompany(d, id)
cs.AddScope(id, []string{"old-domain.com"}, "initial")
// UpdateScope replaces
added, invalid, errs := cs.UpdateScope(id, []string{"new-domain.com"}, "replacement")
if added != 1 || invalid != 0 || len(errs) != 0 {
t.Errorf("UpdateScope: added=%d invalid=%d errs=%v", added, invalid, errs)
}
scope, _ := cs.GetScope(id)
if len(scope) != 1 || scope[0].Domain != "new-domain.com" {
t.Errorf("UpdateScope: expected new-domain.com only, got %+v", scope)
}
// Invalid replacement input must not turn a partial validation response into
// a destructive replacement of the existing rules.
added, invalid, validationErrors, err := cs.UpdateScopeInputsChecked(id, []ScopeInput{
{Kind: "domain", Value: "co.uk"},
}, "invalid replacement")
var validationErr *CompanyScopeValidationError
if added != 0 || invalid != 1 || len(validationErrors) != 1 || !errors.As(err, &validationErr) {
t.Fatalf("invalid replacement: added=%d invalid=%d validation=%v err=%v", added, invalid, validationErrors, err)
}
scope, err = cs.GetScope(id)
if err != nil {
t.Fatal(err)
}
if len(scope) != 1 || scope[0].Domain != "new-domain.com" {
t.Fatalf("invalid replacement changed existing scope: %+v", scope)
}
}
func TestUpdateScopeRollsBackOnInsertFailure(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v)", err)
}
defer d.Close()
cs := d.Companies()
name := fmt.Sprintf("Atomic Scope Update %d", time.Now().UnixNano())
id, _, err := cs.UpsertCompany(name, "")
if err != nil {
t.Fatal(err)
}
defer cleanupCompany(d, id)
oldDomain := fmt.Sprintf("old-%d.example", time.Now().UnixNano())
added, _, invalid, addErrors := cs.AddScopeInputs(id, []ScopeInput{{Kind: "domain", Value: oldDomain}}, "initial")
if added != 1 || invalid != 0 || len(addErrors) != 0 {
t.Fatalf("seed scope: added=%d invalid=%d errors=%v", added, invalid, addErrors)
}
added, invalid, updateErrors := cs.UpdateScopeInputs(id, []ScopeInput{
{Kind: "keyword", Value: "invalid\x00postgres-text"},
}, "replacement")
if added != 0 || invalid != 0 || len(updateErrors) == 0 {
t.Fatalf("failed update result: added=%d invalid=%d errors=%v", added, invalid, updateErrors)
}
scope, err := cs.GetScope(id)
if err != nil {
t.Fatal(err)
}
if len(scope) != 1 || scope[0].Domain != oldDomain {
t.Fatalf("failed replacement did not preserve old scope: %+v", scope)
}
}
func TestDeleteCompany(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v)", err)
}
defer d.Close()
cs := d.Companies()
id, _, err := cs.UpsertCompany("DeleteMeCorp", "")
if err != nil {
t.Fatal(err)
}
cs.AddScope(id, []string{"deletetest.com"}, "test")
if err := cs.DeleteCompany(id); err != nil {
t.Fatal(err)
}
c, err := cs.GetCompany(id)
if err != nil || c != nil {
t.Error("expected company to be gone")
}
// scope should be cascade-deleted
scope, _ := cs.GetScope(id)
if len(scope) != 0 {
t.Errorf("expected scope cascade-deleted, got %d rules", len(scope))
}
}
func TestDeleteCompanyWithAssetsDeletesBoth(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v)", err)
}
defer d.Close()
cs := d.Companies()
stamp := time.Now().UnixNano()
id, _, err := cs.UpsertCompany(fmt.Sprintf("Delete Assets Company %d", stamp), "")
if err != nil {
t.Fatal(err)
}
defer cleanupCompany(d, id)
var assetID int64
domain := fmt.Sprintf("delete-assets-%d.example", stamp)
if err := d.QueryRow(`
INSERT INTO assets(type, domain, root_domain, company_id, company_source)
VALUES ('root_domain', $1, $1, $2, 'explicit')
RETURNING id`, domain, id).Scan(&assetID); err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM assets WHERE id = $1`, assetID) //nolint:errcheck
deleted, err := cs.DeleteCompanyWithAssets(id, true)
if err != nil {
t.Fatal(err)
}
if deleted != 1 {
t.Fatalf("assets deleted=%d want 1", deleted)
}
company, err := cs.GetCompany(id)
if err != nil {
t.Fatal(err)
}
if company != nil {
t.Fatalf("company still exists: %+v", company)
}
var assetsRemaining int
if err := d.QueryRow(`SELECT COUNT(*) FROM assets WHERE id = $1`, assetID).Scan(&assetsRemaining); err != nil {
t.Fatal(err)
}
if assetsRemaining != 0 {
t.Fatalf("asset %d survived company deletion", assetID)
}
}
func TestRecomputeAttribution(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v)", err)
}
defer d.Close()
cs := d.Companies()
as := d.Assets()
id, _, err := cs.UpsertCompany("AttributeTestCorp", "")
if err != nil {
t.Fatal(err)
}
defer cleanupCompany(d, id)
defer d.Exec(`DELETE FROM assets WHERE root_domain = 'attr-test.com'`)
// insert asset before adding scope
assetID, err := as.UpsertRootDomain(UpsertRootDomainReq{Domain: "attr-test.com"})
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM assets WHERE id = $1`, assetID)
// asset should not be attributed yet
var companyID *int64
d.QueryRow(`SELECT company_id FROM assets WHERE id = $1`, assetID).Scan(&companyID)
if companyID != nil {
t.Error("expected no company before scope added")
}
// add scope and recompute
cs.AddScope(id, []string{"attr-test.com"}, "test")
if err := cs.RecomputeAttribution(); err != nil {
t.Fatal(err)
}
d.QueryRow(`SELECT company_id FROM assets WHERE id = $1`, assetID).Scan(&companyID)
if companyID == nil || *companyID != id {
t.Errorf("RecomputeAttribution: expected company %d, got %v", id, companyID)
}
}
func TestListCompanies(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v)", err)
}
defer d.Close()
cs := d.Companies()
id, _, err := cs.UpsertCompany("ListTestCorp", "")
if err != nil {
t.Fatal(err)
}
defer cleanupCompany(d, id)
companies, err := cs.ListCompanies()
if err != nil {
t.Fatal(err)
}
found := false
for _, c := range companies {
if c.ID == id {
found = true
}
}
if !found {
t.Error("ListCompanies: created company not found")
}
}
+15
View File
@@ -0,0 +1,15 @@
package db
import "testing"
func TestCompanyScopeMutationLockDoesNotReuseInfrastructureLocks(t *testing.T) {
reserved := map[string]int64{
"schema migration": 7337741001,
"cross-package test suite": 7337741002,
}
for name, key := range reserved {
if companyScopeMutationLock == key {
t.Fatalf("company scope mutation lock reuses the %s advisory lock key %d", name, key)
}
}
}
+272
View File
@@ -0,0 +1,272 @@
package db
import (
"fmt"
"net"
"net/url"
"strings"
"unicode"
"golang.org/x/net/idna"
"golang.org/x/net/publicsuffix"
)
// ParsedScope is one parsed asset-scope entry. Its kind selects Domain, Net, or Value.
// Used internally by the company scope parsers and CompanyStore.
type ParsedScope struct {
Kind string // "domain" | "ip" | "cidr" | "icp" | "keyword"
Domain string // normalized registrable/root domain (kind=domain)
Net string // normalized CIDR, single IP as /32 or /128 (kind=ip|cidr)
Value string // normalized text (kind=icp|keyword)
Raw string // original input line
}
// ScopeInput is the structured API form for a company scope rule. Empty Kind
// uses the same automatic classification as the single-textarea UI.
type ScopeInput struct {
Kind string `json:"kind,omitempty"`
Value string `json:"value"`
}
// NormalizeICP removes every Unicode whitespace character and folds case. ICP
// matching intentionally performs no fuzzy or punctuation normalization.
func NormalizeICP(value string) string {
return strings.ToLower(strings.Map(func(r rune) rune {
if unicode.IsSpace(r) {
return -1
}
return r
}, strings.TrimSpace(value)))
}
func normalizeKeyword(value string) string {
return strings.ToLower(strings.Join(strings.Fields(value), " "))
}
func looksLikeIPAddress(value string) bool {
value = strings.TrimSpace(value)
if strings.Contains(value, "://") || strings.IndexFunc(value, unicode.IsSpace) >= 0 {
return false
}
if strings.Count(value, ":") >= 2 {
// Require an IPv6-looking prefix. This still catches malformed values such
// as 2001:db8::zz without treating ordinary colon-delimited keywords as IPs.
parts := strings.Split(value, ":")
validSegments := 0
for _, part := range parts {
if part == "" {
if validSegments > 0 || strings.HasPrefix(value, "::") {
return true
}
continue
}
if len(part) > 4 {
return false
}
for _, r := range part {
if !((r >= '0' && r <= '9') || (r >= 'a' && r <= 'f') || (r >= 'A' && r <= 'F')) {
return false
}
}
validSegments++
if validSegments >= 2 {
return true
}
}
return false
}
if !strings.Contains(value, ".") {
return false
}
for _, r := range value {
if r != '.' && (r < '0' || r > '9') {
return false
}
}
return true
}
// ParseScopeInput validates an explicitly typed rule. Legacy callers can omit
// Kind and use the same automatic classification as the single-textarea UI.
func ParseScopeInput(input ScopeInput) (ParsedScope, error) {
kind := strings.ToLower(strings.TrimSpace(input.Kind))
raw := strings.TrimSpace(input.Value)
if kind == "" {
return ParseAutoScopeLine(raw)
}
switch kind {
case "domain", "ip", "cidr":
rule, err := ParseScopeLine(raw)
if err != nil {
return rule, err
}
if rule.Kind != kind {
return ParsedScope{Kind: kind, Raw: raw}, fmt.Errorf("%q 不是有效的 %s 范围", raw, kind)
}
return rule, nil
case "icp":
value := NormalizeICP(raw)
if value == "" {
return ParsedScope{Kind: kind, Raw: raw}, fmt.Errorf("ICP 不能为空")
}
return ParsedScope{Kind: kind, Value: value, Raw: raw}, nil
case "keyword":
value := normalizeKeyword(raw)
if value == "" {
return ParsedScope{Kind: kind, Raw: raw}, fmt.Errorf("企业关键词不能为空")
}
return ParsedScope{Kind: kind, Value: value, Raw: raw}, nil
default:
return ParsedScope{Kind: kind, Raw: raw}, fmt.Errorf("不支持的范围类型: %s", kind)
}
}
// ParseAutoScopeLine classifies one untyped textarea line. Network-looking and
// domain-looking values remain strict so malformed ranges do not silently become
// Agent keywords; all other non-empty text is a keyword.
func ParseAutoScopeLine(line string) (ParsedScope, error) {
raw := strings.TrimSpace(line)
if raw == "" {
return ParsedScope{}, fmt.Errorf("空行")
}
if _, _, err := net.ParseCIDR(raw); err == nil {
return ParseScopeLine(raw)
}
if ip := net.ParseIP(raw); ip != nil {
return ParseScopeLine(raw)
}
if slash := strings.LastIndexByte(raw, '/'); slash > 0 {
address := strings.TrimSpace(raw[:slash])
if net.ParseIP(address) != nil || looksLikeIPAddress(address) {
return ParsedScope{Raw: raw}, fmt.Errorf("无效 CIDR: %s", raw)
}
}
if looksLikeIPAddress(raw) {
return ParsedScope{Raw: raw}, fmt.Errorf("无效 IP: %s", raw)
}
looksLikeDomain := strings.Contains(raw, "://") ||
(strings.Contains(raw, ".") && strings.IndexFunc(raw, unicode.IsSpace) < 0)
if looksLikeDomain {
return ParseScopeLine(raw)
}
// 备案号本身不含点号(如 京ICP备12345678号-1)。带点的文本多半掺了域名或版本号,
// 按 ICP 存下来只会得到一条永远匹配不上任何资产的死规则 —— ICP 归属走的是精确
// 相等比较(见 companies.go 的 kind='icp' 归属查询),所以这类文本归为关键词。
lower := strings.ToLower(raw)
if !strings.ContainsAny(raw, "..。") &&
(strings.Contains(lower, "icp") || strings.Contains(raw, "备案")) {
return ParseScopeInput(ScopeInput{Kind: "icp", Value: raw})
}
return ParseScopeInput(ScopeInput{Kind: "keyword", Value: raw})
}
func scopeHostname(raw string) (string, error) {
candidate := strings.TrimSpace(raw)
if candidate == "" {
return "", fmt.Errorf("主机名为空")
}
if strings.HasPrefix(candidate, "//") {
candidate = "http:" + candidate
} else if !strings.Contains(candidate, "://") {
candidate = "http://" + candidate
}
parsed, err := url.Parse(candidate)
if err != nil || parsed.Host == "" {
if err == nil {
err = fmt.Errorf("缺少主机名")
}
return "", err
}
host := strings.TrimSuffix(strings.TrimSpace(parsed.Hostname()), ".")
if host == "" {
return "", fmt.Errorf("主机名为空")
}
if ip := net.ParseIP(host); ip != nil {
return ip.String(), nil
}
host, err = idna.Lookup.ToASCII(host)
if err != nil {
return "", err
}
host = strings.ToLower(host)
if len(host) > 253 {
return "", fmt.Errorf("域名超过 253 个字符")
}
labels := strings.Split(host, ".")
if len(labels) < 2 {
return "", fmt.Errorf("域名至少需要两个标签")
}
for _, label := range labels {
if label == "" || len(label) > 63 || label[0] == '-' || label[len(label)-1] == '-' {
return "", fmt.Errorf("域名标签无效")
}
for _, r := range label {
if !((r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') || r == '-') {
return "", fmt.Errorf("域名包含无效字符")
}
}
}
return host, nil
}
// ParseScopeLine classifies and validates one scope line (root domain / IP /
// CIDR). Guardrails reject bare TLDs and over-broad networks so a rule can never
// swallow the internet. IP ranges must be expressed as CIDR.
func ParseScopeLine(line string) (ParsedScope, error) {
raw := strings.TrimSpace(line)
r := ParsedScope{Raw: raw}
if raw == "" {
return r, fmt.Errorf("空行")
}
// CIDR first because URL parsing treats its slash as a path separator.
if _, ipnet, err := net.ParseCIDR(raw); err == nil {
ones, bits := ipnet.Mask.Size()
if bits == 32 && ones < 16 {
return r, fmt.Errorf("网段过宽(IPv4 需 >= /16): %s", raw)
}
if bits == 128 && ones < 32 {
return r, fmt.Errorf("网段过宽(IPv6 需 >= /32): %s", raw)
}
r.Kind, r.Net = "cidr", ipnet.String()
return r, nil
}
// Single IP.
if ip := net.ParseIP(raw); ip != nil {
r.Kind = "ip"
if ip.To4() != nil {
r.Net = ip.String() + "/32"
} else {
r.Net = ip.String() + "/128"
}
return r, nil
}
host, err := scopeHostname(raw)
if err != nil {
return r, fmt.Errorf("无法识别为有效域名/IP/CIDR: %s", raw)
}
if ip := net.ParseIP(host); ip != nil {
r.Kind = "ip"
if ip.To4() != nil {
r.Net = ip.String() + "/32"
} else {
r.Net = ip.String() + "/128"
}
return r, nil
}
if looksLikeIPAddress(host) {
return r, fmt.Errorf("无效 IP: %s", raw)
}
if strings.Contains(raw, "-") && strings.Count(raw, ".") >= 6 {
return r, fmt.Errorf("IP 段请用 CIDR 表示(如 1.2.3.0/24): %s", raw)
}
// Domain (registrable). Reject bare TLDs / public suffixes.
d := DomainKey(host)
if suf, icann := publicsuffix.PublicSuffix(d); icann && suf == d {
return r, fmt.Errorf("不能用裸 TLD 作为范围: %s", raw)
}
r.Kind, r.Domain = "domain", d
return r, nil
}
+250
View File
@@ -0,0 +1,250 @@
package db
import (
"errors"
"fmt"
"strings"
"testing"
"time"
)
func TestAssetUpsertWaitsForCompanyScopeMutation(t *testing.T) {
d, assets, companies := testSetup(t)
defer d.Close()
stamp := time.Now().UnixNano()
domain := fmt.Sprintf("scope-lock-%d.example", stamp)
oldCompany, _, err := companies.UpsertCompany(fmt.Sprintf("Scope Lock Old %d", stamp), "")
if err != nil {
t.Fatal(err)
}
newCompany, _, err := companies.UpsertCompany(fmt.Sprintf("Scope Lock New %d", stamp), "")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
_, _ = d.Exec(`DELETE FROM assets WHERE domain = $1`, domain)
_, _ = d.Exec(`DELETE FROM companies WHERE id IN ($1,$2)`, oldCompany, newCompany)
})
if added, _, invalid, validationErrors, err := companies.AddScopeInputsChecked(oldCompany, []ScopeInput{
{Kind: "domain", Value: domain},
}, "test"); err != nil || added != 1 || invalid != 0 || len(validationErrors) != 0 {
t.Fatalf("seed old scope: added=%d invalid=%d validation=%v err=%v", added, invalid, validationErrors, err)
}
mutation, err := d.Begin()
if err != nil {
t.Fatal(err)
}
defer mutation.Rollback() //nolint:errcheck
if err := lockCompanyScopeMutation(mutation); err != nil {
t.Fatal(err)
}
if _, err := mutation.Exec(`DELETE FROM company_scope WHERE company_id=$1 AND kind='domain' AND domain=$2`, oldCompany, domain); err != nil {
t.Fatal(err)
}
if _, err := insertScopeRuleTx(mutation, newCompany, ParsedScope{
Kind: "domain", Domain: domain, Raw: domain,
}, "test"); err != nil {
t.Fatal(err)
}
started := make(chan struct{})
type result struct {
id int64
err error
}
resultCh := make(chan result, 1)
go func() {
close(started)
id, err := assets.UpsertRootDomain(UpsertRootDomainReq{Domain: domain})
resultCh <- result{id: id, err: err}
}()
<-started
select {
case got := <-resultCh:
t.Fatalf("asset upsert escaped scope mutation lock: %+v", got)
case <-time.After(100 * time.Millisecond):
}
if err := mutation.Commit(); err != nil {
t.Fatal(err)
}
var got result
select {
case got = <-resultCh:
case <-time.After(3 * time.Second):
t.Fatal("asset upsert did not resume after scope mutation committed")
}
if got.err != nil {
t.Fatal(got.err)
}
var companyID *int64
if err := d.QueryRow(`SELECT company_id FROM assets WHERE id=$1`, got.id).Scan(&companyID); err != nil {
t.Fatal(err)
}
if companyID == nil || *companyID != newCompany {
t.Fatalf("asset retained stale company: got=%v want=%d", companyID, newCompany)
}
}
func TestResolveAndRecomputeUseStableCompanyIDTieBreak(t *testing.T) {
d, assets, companies := testSetup(t)
defer d.Close()
stamp := time.Now().UnixNano()
first, _, err := companies.UpsertCompany(fmt.Sprintf("Tie Break First %d", stamp), "")
if err != nil {
t.Fatal(err)
}
second, _, err := companies.UpsertCompany(fmt.Sprintf("Tie Break Second %d", stamp), "")
if err != nil {
t.Fatal(err)
}
want := min(first, second)
domain := fmt.Sprintf("tie-break-%d.example", stamp)
segment := uint64(stamp) & 0xffff
network := fmt.Sprintf("2001:db8:%x:1::/64", segment)
ip := fmt.Sprintf("2001:db8:%x:1::42", segment)
icp := fmt.Sprintf("ICP-TIE-%d", stamp)
appName := fmt.Sprintf("Tie Break App %d", stamp)
t.Cleanup(func() {
_, _ = d.Exec(`DELETE FROM assets WHERE domain=$1 OR ip=$2 OR app_name=$3`, domain, ip, appName)
_, _ = d.Exec(`DELETE FROM companies WHERE id IN ($1,$2)`, first, second)
})
rules := []ScopeInput{
{Kind: "domain", Value: domain},
{Kind: "cidr", Value: network},
{Kind: "icp", Value: icp},
}
for _, companyID := range []int64{second, first} {
if added, _, invalid, validationErrors, err := companies.AddScopeInputsChecked(companyID, rules, "test"); err != nil || added != len(rules) || invalid != 0 || len(validationErrors) != 0 {
t.Fatalf("add company %d rules: added=%d invalid=%d validation=%v err=%v", companyID, added, invalid, validationErrors, err)
}
}
assertResolved := func(label string, got *int64, err error) {
t.Helper()
if err != nil {
t.Fatalf("%s resolve: %v", label, err)
}
if got == nil || *got != want {
t.Fatalf("%s resolve=%v want lowest company id %d", label, got, want)
}
}
got, err := companies.ResolveCompany(domain, "")
assertResolved("domain", got, err)
got, err = companies.ResolveCompany("", ip)
assertResolved("network", got, err)
got, err = companies.ResolveCompanyWithICP("", "", icp)
assertResolved("icp", got, err)
rootID, err := assets.UpsertRootDomain(UpsertRootDomainReq{Domain: domain})
if err != nil {
t.Fatal(err)
}
ipID, err := assets.UpsertIP(UpsertIPReq{IP: ip})
if err != nil {
t.Fatal(err)
}
appID, err := assets.UpsertApp(UpsertAppReq{Name: appName, ICP: icp})
if err != nil {
t.Fatal(err)
}
assertAssets := func(stage string) {
t.Helper()
for _, id := range []int64{rootID, ipID, appID} {
var got *int64
if err := d.QueryRow(`SELECT company_id FROM assets WHERE id=$1`, id).Scan(&got); err != nil {
t.Fatal(err)
}
if got == nil || *got != want {
t.Fatalf("%s asset %d company=%v want=%d", stage, id, got, want)
}
}
}
assertAssets("live")
if err := companies.RecomputeAttribution(); err != nil {
t.Fatal(err)
}
assertAssets("recomputed")
}
func TestCompanyScopeLimitsAndCheckedErrors(t *testing.T) {
boundary := strings.Repeat("界", MaxCompanyScopeRawRunes)
if err := ValidateCompanyScopeInputBounds([]ScopeInput{{Kind: "keyword", Value: boundary}}); err != nil {
t.Fatalf("exact raw rune boundary rejected: %v", err)
}
var validationErr *CompanyScopeValidationError
if err := ValidateCompanyScopeInputBounds([]ScopeInput{{Kind: "keyword", Value: boundary + "界"}}); !errors.As(err, &validationErr) {
t.Fatalf("oversized raw value error=%v want CompanyScopeValidationError", err)
}
// 条数不再设上限,只校验单条长度。
if err := ValidateCompanyScopeInputBounds(make([]ScopeInput, 1000)); err != nil {
t.Fatalf("rule count should be unbounded, got %v", err)
}
d, _, companies := testSetup(t)
defer d.Close()
stamp := time.Now().UnixNano()
companyID, _, err := companies.UpsertCompany(fmt.Sprintf("Scope Limits %d", stamp), "")
if err != nil {
t.Fatal(err)
}
defer cleanupCompany(d, companyID)
// 曾经封顶 256 条,逐个 IP / 域名录范围的企业很容易撞上;现在不限条数。
const bulk = 300
rules := make([]ScopeInput, bulk)
for i := range rules {
rules[i] = ScopeInput{Kind: "keyword", Value: fmt.Sprintf("limit-%d-%d", stamp, i)}
}
if added, _, _, _, err := companies.AddScopeInputsChecked(companyID, rules, "test"); err != nil || added != bulk {
t.Fatalf("bulk scope insert: added=%d want=%d err=%v", added, bulk, err)
}
if added, _, _, _, err := companies.AddScopeInputsChecked(companyID, []ScopeInput{
{Kind: "keyword", Value: fmt.Sprintf("limit-%d-extra", stamp)},
}, "test"); added != 1 || err != nil {
t.Fatalf("append past the old cap: added=%d err=%v", added, err)
}
var count int
if err := d.QueryRow(`SELECT count(*) FROM company_scope WHERE company_id=$1`, companyID).Scan(&count); err != nil {
t.Fatal(err)
}
if count != bulk+1 {
t.Fatalf("scope count=%d want=%d", count, bulk+1)
}
missingID := companyID + 1_000_000_000
if _, _, _, _, err := companies.AddScopeInputsChecked(missingID, nil, "test"); !errors.Is(err, ErrCompanyNotFound) {
t.Fatalf("missing add error=%v want ErrCompanyNotFound", err)
}
if _, _, _, err := companies.UpdateScopeInputsChecked(missingID, nil, "test"); !errors.Is(err, ErrCompanyNotFound) {
t.Fatalf("missing update error=%v want ErrCompanyNotFound", err)
}
if _, err := companies.DeleteCompanyWithAssets(missingID, true); !errors.Is(err, ErrCompanyNotFound) {
t.Fatalf("missing delete error=%v want ErrCompanyNotFound", err)
}
}
func TestCompanyScopeCheckedReturnsSystemErrorSeparately(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v)", err)
}
companies := d.Companies()
if err := d.Close(); err != nil {
t.Fatal(err)
}
_, _, _, validationErrors, err := companies.AddScopeInputsChecked(1, nil, "test")
if err == nil {
t.Fatal("closed database did not return a system error")
}
if len(validationErrors) != 0 {
t.Fatalf("system error leaked into validation errors: %v", validationErrors)
}
var validationErr *CompanyScopeValidationError
if errors.As(err, &validationErr) || errors.Is(err, ErrCompanyNotFound) {
t.Fatalf("system error misclassified: %v", err)
}
}
+405
View File
@@ -0,0 +1,405 @@
package db
import (
"fmt"
"testing"
"time"
)
// TestParseScopeLine covers classification + guardrails without a DB.
func TestParseScopeLine(t *testing.T) {
cases := []struct {
in string
kind string
wantErr bool
}{
{"example.com", "domain", false},
{"https://sub.example.com/path", "domain", false},
{"1.2.3.4", "ip", false},
{"10.0.0.0/8", "", true}, // over-broad IPv4 (< /16)
{"198.51.100.0/24", "cidr", false},
{"co.uk", "", true}, // bare public suffix
{"not a host", "", true},
{"1.2.3.1-1.2.3.9", "", true}, // ranges must be CIDR
}
for _, c := range cases {
r, err := ParseScopeLine(c.in)
if c.wantErr {
if err == nil {
t.Errorf("ParseScopeLine(%q) want error, got %+v", c.in, r)
}
continue
}
if err != nil {
t.Errorf("ParseScopeLine(%q) unexpected error: %v", c.in, err)
continue
}
if r.Kind != c.kind {
t.Errorf("ParseScopeLine(%q) kind=%q want %q", c.in, r.Kind, c.kind)
}
}
}
func TestParseAutoScopeLine(t *testing.T) {
cases := []struct {
name string
input string
kind string
normalized string
wantErr bool
}{
{name: "domain", input: "example.com", kind: "domain", normalized: "example.com"},
{name: "url", input: "https://sub.example.com/path", kind: "domain", normalized: "sub.example.com"},
{name: "url query", input: "https://example.com/path?source=x", kind: "domain", normalized: "example.com"},
{name: "url credentials", input: "https://user:pass@example.com/path", kind: "domain", normalized: "example.com"},
{name: "url ip", input: "http://203.0.113.10/path", kind: "ip", normalized: "203.0.113.10/32"},
{name: "ipv4", input: "203.0.113.10", kind: "ip", normalized: "203.0.113.10/32"},
{name: "ipv6", input: "2001:db8::10", kind: "ip", normalized: "2001:db8::10/128"},
{name: "cidr", input: "198.51.100.0/24", kind: "cidr", normalized: "198.51.100.0/24"},
{name: "icp latin", input: "京 ICP备 123号", kind: "icp", normalized: "京icp备123号"},
{name: "icp chinese", input: "沪网备案 9988", kind: "icp", normalized: "沪网备案9988"},
{name: "icp domain", input: "icp.example.com", kind: "domain", normalized: "icp.example.com"},
{name: "icp url query", input: "https://example.com/path?icp=1", kind: "domain", normalized: "example.com"},
// 备案号不含点号:掺了域名/版本号的描述性文字归关键词,否则会存成一条
// 永远匹配不上的死 ICP 规则。
{name: "icp with domain text", input: "备案 www.example.com", kind: "keyword", normalized: "备案 www.example.com"},
{name: "icp with version text", input: "某公司 ICP v1.0", kind: "keyword", normalized: "某公司 icp v1.0"},
{name: "icp fullwidth dot", input: "备案 例.com", kind: "keyword", normalized: "备案 例.com"},
{name: "keyword", input: " ACME Security ", kind: "keyword", normalized: "acme security"},
{name: "colon keyword", input: "ACME: Cloud: Security", kind: "keyword", normalized: "acme: cloud: security"},
{name: "empty", input: " ", wantErr: true},
{name: "invalid cidr", input: "10.0.0.0/not-a-prefix", wantErr: true},
{name: "invalid ipv4", input: "999.0.0.1", wantErr: true},
{name: "invalid ipv6", input: "2001:db8::zz", wantErr: true},
{name: "invalid ipv6 cidr", input: "2001:db8::zz/64", wantErr: true},
{name: "overbroad cidr", input: "10.0.0.0/8", wantErr: true},
{name: "bare suffix", input: "co.uk", wantErr: true},
{name: "empty domain label", input: "foo..example.com", wantErr: true},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
rule, err := ParseAutoScopeLine(tc.input)
if tc.wantErr {
if err == nil {
t.Fatalf("ParseAutoScopeLine(%q) = %+v, want error", tc.input, rule)
}
return
}
if err != nil {
t.Fatalf("ParseAutoScopeLine(%q): %v", tc.input, err)
}
if rule.Kind != tc.kind {
t.Fatalf("kind=%q want %q", rule.Kind, tc.kind)
}
got := rule.Value
if rule.Kind == "domain" {
got = rule.Domain
} else if rule.Kind == "ip" || rule.Kind == "cidr" {
got = rule.Net
}
if got != tc.normalized {
t.Fatalf("normalized=%q want %q", got, tc.normalized)
}
})
}
}
func TestExplicitCompanyAttributionSurvivesScopeRebuild(t *testing.T) {
d, as, cs := testSetup(t)
defer d.Close()
stamp := time.Now().UnixNano()
explicitCompany, _, err := cs.UpsertCompany(fmt.Sprintf("Explicit Attribution %d", stamp), "")
if err != nil {
t.Fatal(err)
}
autoCompany, _, err := cs.UpsertCompany(fmt.Sprintf("Automatic Attribution %d", stamp), "")
if err != nil {
t.Fatal(err)
}
defer cleanupCompany(d, explicitCompany)
defer cleanupCompany(d, autoCompany)
domain := fmt.Sprintf("explicit-%d.invalid", stamp)
icp := fmt.Sprintf("ICP-%d", stamp)
network := fmt.Sprintf("2001:db8:%x::/64", uint64(stamp)&0xffff)
ip := fmt.Sprintf("2001:db8:%x::10", uint64(stamp)&0xffff)
// A pre-existing row with company_id and no provenance value represents old
// installations. The schema default conservatively treats it as explicit.
var assetID int64
if err := d.QueryRow(`INSERT INTO assets(type,domain,root_domain,company_id)
VALUES ('root_domain',$1,$1,$2) RETURNING id`, domain, explicitCompany).Scan(&assetID); err != nil {
t.Fatal(err)
}
defer deleteAsset(d, assetID)
appID, err := as.UpsertApp(UpsertAppReq{
Name: fmt.Sprintf("explicit-app-%d", stamp), ICP: icp, CompanyID: &explicitCompany,
})
if err != nil {
t.Fatal(err)
}
defer deleteAsset(d, appID)
autoAssetID, err := as.UpsertIP(UpsertIPReq{IP: ip})
if err != nil {
t.Fatal(err)
}
defer deleteAsset(d, autoAssetID)
rules := []ScopeInput{
{Kind: "domain", Value: domain},
{Kind: "icp", Value: icp},
{Kind: "cidr", Value: network},
}
if added, _, invalid, errs := cs.AddScopeInputs(autoCompany, rules, "test"); added != len(rules) || invalid != 0 {
t.Fatalf("AddScopeInputs: added=%d invalid=%d errors=%v", added, invalid, errs)
}
assertCompany := func(id, want int64, source string) {
t.Helper()
var got *int64
var gotSource string
if err := d.QueryRow(`SELECT company_id,company_source FROM assets WHERE id=$1`, id).Scan(&got, &gotSource); err != nil {
t.Fatal(err)
}
if got == nil || *got != want || gotSource != source {
t.Fatalf("asset %d company=%v source=%q, want %d/%q", id, got, gotSource, want, source)
}
}
assertCompany(assetID, explicitCompany, "explicit")
assertCompany(appID, explicitCompany, "explicit")
assertCompany(autoAssetID, autoCompany, "scope")
// Replacing every matching rule detaches only the automatically-owned row.
if _, invalid, errs := cs.UpdateScopeInputs(autoCompany, []ScopeInput{
{Kind: "domain", Value: fmt.Sprintf("replacement-%d.invalid", stamp)},
}, "test"); invalid != 0 || len(errs) != 0 {
t.Fatalf("UpdateScopeInputs: invalid=%d errors=%v", invalid, errs)
}
assertCompany(assetID, explicitCompany, "explicit")
assertCompany(appID, explicitCompany, "explicit")
var autoCompanyID *int64
var autoSource string
if err := d.QueryRow(`SELECT company_id,company_source FROM assets WHERE id=$1`, autoAssetID).Scan(&autoCompanyID, &autoSource); err != nil {
t.Fatal(err)
}
if autoCompanyID != nil || autoSource != "scope" {
t.Fatalf("automatic asset was not detached: company=%v source=%q", autoCompanyID, autoSource)
}
if err := cs.RecomputeAttribution(); err != nil {
t.Fatal(err)
}
assertCompany(assetID, explicitCompany, "explicit")
assertCompany(appID, explicitCompany, "explicit")
// Deleting the explicitly selected company detaches through the FK, then the
// transactional rebuild may adopt the assets into a still-valid scope.
if _, invalid, errs := cs.UpdateScopeInputs(autoCompany, []ScopeInput{
{Kind: "domain", Value: domain},
{Kind: "icp", Value: icp},
}, "test"); invalid != 0 || len(errs) != 0 {
t.Fatalf("restore fallback scope: invalid=%d errors=%v", invalid, errs)
}
if err := cs.DeleteCompany(explicitCompany); err != nil {
t.Fatal(err)
}
assertCompany(assetID, autoCompany, "scope")
assertCompany(appID, autoCompany, "scope")
}
func TestParseStructuredCompanyScope(t *testing.T) {
icp, err := ParseScopeInput(ScopeInput{Kind: "ICP", Value: " 京ICP 备 123号-1\t"})
if err != nil {
t.Fatalf("parse ICP: %v", err)
}
if icp.Kind != "icp" || icp.Value != "京icp备123号-1" {
t.Fatalf("unexpected normalized ICP: %+v", icp)
}
keyword, err := ParseScopeInput(ScopeInput{Kind: "keyword", Value: " ACME Security "})
if err != nil {
t.Fatalf("parse keyword: %v", err)
}
if keyword.Value != "acme security" {
t.Fatalf("unexpected normalized keyword: %+v", keyword)
}
if _, err := ParseScopeInput(ScopeInput{Kind: "ip", Value: "example.com"}); err == nil {
t.Fatal("typed IP accepted a domain")
}
}
func TestCompanyICPAttribution(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
// 关连接必须走 t.Cleanup 且**注册在清理之前**:t.Cleanup 是后进先出,
// 先注册关闭 → 关闭最后执行,下面的数据清理才连得上库。
// 原先这里是 `defer d.Close()`:defer 在函数返回时先跑,t.Cleanup 在那之后
// 才执行,于是清理语句全落在**已关闭的连接**上、错误又被 `_, _ =` 丢弃,
// 资产与公司就永久残留在库里。残留本身不会立刻报错,但本用例用
// `MAX(companies.id)+1` 当假 TaskID 给资产打标(见下方 suffix),
// 一旦这个数字与别的用例的任务 id 撞上,那个用例按「恰好 N 个资产」的断言
// 就会莫名失败——排查成本极高。
t.Cleanup(func() { d.Close() })
var suffix int64
if err := d.QueryRow(`SELECT COALESCE(MAX(id),0)+1 FROM companies`).Scan(&suffix); err != nil {
t.Fatal(err)
}
cs := d.Companies()
as := d.Assets()
companyID, _, err := cs.UpsertCompany(fmt.Sprintf("ICP Scope Co %d", suffix), "")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
// 不吞错误:清理失败会污染后续用例,必须让它在本次运行里显形。
if _, err := d.Exec(`DELETE FROM assets WHERE task_ids @> ARRAY[$1]::bigint[]`, suffix); err != nil {
t.Errorf("清理测试资产失败: %v", err)
}
if _, err := d.Exec(`DELETE FROM companies WHERE id=$1`, companyID); err != nil {
t.Errorf("清理测试公司失败: %v", err)
}
})
// A keyword can guide an Agent, but must never claim an asset by its name.
added, _, invalid, errs := cs.AddScopeInputs(companyID, []ScopeInput{
{Kind: "icp", Value: "京 ICP备 998877号"},
{Kind: "keyword", Value: "ICP Scope"},
}, "unit test")
if added != 2 || invalid != 0 || len(errs) != 0 {
t.Fatalf("add structured scope: added=%d invalid=%d errors=%v", added, invalid, errs)
}
rootID, err := as.UpsertRootDomain(UpsertRootDomainReq{
Domain: fmt.Sprintf("icp-scope-%d.example", suffix), ICP: "京icp备998877号", TaskID: suffix,
})
if err != nil {
t.Fatal(err)
}
appID, err := as.UpsertApp(UpsertAppReq{
Name: fmt.Sprintf("ICP Scope Keyword Only %d", suffix), TaskID: suffix,
})
if err != nil {
t.Fatal(err)
}
icpAppID, err := as.UpsertApp(UpsertAppReq{
Name: fmt.Sprintf("ICP Matched App %d", suffix), ICP: " 京 ICP备 998877号 ", TaskID: suffix,
})
if err != nil {
t.Fatal(err)
}
assertCompany := func(assetID int64, want *int64) {
t.Helper()
var got *int64
if err := d.QueryRow(`SELECT company_id FROM assets WHERE id=$1`, assetID).Scan(&got); err != nil {
t.Fatal(err)
}
if want == nil && got != nil {
t.Fatalf("asset %d attributed by keyword: %d", assetID, *got)
}
if want != nil && (got == nil || *got != *want) {
t.Fatalf("asset %d company=%v want %d", assetID, got, *want)
}
}
assertCompany(rootID, &companyID)
assertCompany(appID, nil)
assertCompany(icpAppID, &companyID)
}
// TestCompanyScopeAttribution exercises the full loop against dev PG: create
// company (unique name), add scope, and verify auto-attribution at insert time,
// backfill of a pre-existing asset, CIDR + domain-suffix matching, and that
// out-of-scope assets stay unattributed.
func TestCompanyScopeAttribution(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
as := d.Assets()
cs := d.Companies()
var startMax int64
if err := d.QueryRow(`SELECT COALESCE(MAX(id),0) FROM assets`).Scan(&startMax); err != nil {
d.Close()
t.Fatalf("startMax: %v", err)
}
t.Cleanup(func() {
_, _ = d.Exec(`DELETE FROM assets WHERE id > $1`, startMax)
d.Close()
})
uniq := startMax + 1
root := fmt.Sprintf("scopetest%d.com", uniq)
sub := "api." + root
ipIn := "198.51.100.9"
ipOut := "203.0.113.9"
outDomain := fmt.Sprintf("other%d.net", uniq)
cid, _, err := cs.UpsertCompany(fmt.Sprintf("ScopeCo %d", uniq), "")
if err != nil {
t.Fatalf("UpsertCompany: %v", err)
}
// a pre-existing asset (inserted BEFORE any scope) — must be back-filled.
preID, err := as.UpsertSubdomain(UpsertSubdomainReq{Domain: sub})
if err != nil {
t.Fatalf("pre upsert: %v", err)
}
var preCompanyID *int64
d.QueryRow(`SELECT company_id FROM assets WHERE id = $1`, preID).Scan(&preCompanyID)
if preCompanyID != nil {
t.Fatalf("pre-scope asset should be unattributed, got %v", *preCompanyID)
}
cs.AddScope(cid, []string{root, "198.51.100.0/24"}, "unit test")
mustCid := func(id int64, want int64, label string) {
var cID *int64
d.QueryRow(`SELECT company_id FROM assets WHERE id = $1`, id).Scan(&cID)
if cID == nil {
t.Fatalf("%s company_id = nil, want %d", label, want)
}
if *cID != want {
t.Fatalf("%s company_id = %d, want %d", label, *cID, want)
}
}
mustNil := func(id int64, label string) {
var cID *int64
d.QueryRow(`SELECT company_id FROM assets WHERE id = $1`, id).Scan(&cID)
if cID != nil {
t.Fatalf("%s should be unattributed, got %d", label, *cID)
}
}
// backfill attributed the pre-existing subdomain (domain suffix match).
mustCid(preID, cid, "pre-existing subdomain (backfill)")
// insert-time attribution: ip in CIDR, another subdomain.
ipInID, err := as.UpsertIP(UpsertIPReq{IP: ipIn})
if err != nil {
t.Fatalf("UpsertIP in: %v", err)
}
mustCid(ipInID, cid, "in-CIDR ip (insert-time)")
sub2ID, err := as.UpsertSubdomain(UpsertSubdomainReq{Domain: "www." + root})
if err != nil {
t.Fatalf("UpsertSubdomain: %v", err)
}
mustCid(sub2ID, cid, "new subdomain (insert-time)")
// out of scope stays unattributed.
outID, err := as.UpsertRootDomain(UpsertRootDomainReq{Domain: outDomain})
if err != nil {
t.Fatalf("UpsertRootDomain out: %v", err)
}
mustNil(outID, "out-of-scope domain")
ipOutID, err := as.UpsertIP(UpsertIPReq{IP: ipOut})
if err != nil {
t.Fatalf("UpsertIP out: %v", err)
}
mustNil(ipOutID, "out-of-CIDR ip")
}
+1087
View File
File diff suppressed because it is too large Load Diff
+100
View File
@@ -0,0 +1,100 @@
package db
import "testing"
// TestProfileMaxTokensRoundTrip pins the output-cap columns through the whole
// read/write surface: both column lists (list vs keyed loads) must carry them,
// and an update must not drop them. A miscounted column list silently shifts
// every later Scan target, so this is the test that catches it.
func TestProfileMaxTokensRoundTrip(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
// Close via Cleanup, registered FIRST so it runs LAST: cleanups are LIFO, and a
// plain `defer d.Close()` would fire before them — the row-deleting cleanups
// would then run against a closed pool and silently leave test rows behind.
t.Cleanup(func() { d.Close() })
id, err := d.SaveProfile(&LLMProfile{
Name: "t-maxtok", Format: "openai", Model: "m", APIKey: "k",
MaxTokens: 8192, MaxTokensField: "max_completion_tokens",
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { d.Exec(`DELETE FROM llm_profiles WHERE id=$1`, id) })
// Keyed single-row load (profileColsKey).
p, err := d.ProfileByID(id)
if err != nil || p == nil {
t.Fatalf("ProfileByID: %v, p=%v", err, p)
}
if p.MaxTokens != 8192 || p.MaxTokensField != "max_completion_tokens" {
t.Fatalf("keyed load: max_tokens=%d field=%q", p.MaxTokens, p.MaxTokensField)
}
// List load (profileCols — the hint variant, a separate column list).
ps, err := d.ListProfiles()
if err != nil {
t.Fatal(err)
}
var found *LLMProfile
for _, x := range ps {
if x.ID == id {
found = x
}
}
if found == nil {
t.Fatal("profile missing from ListProfiles")
}
if found.MaxTokens != 8192 || found.MaxTokensField != "max_completion_tokens" {
t.Fatalf("list load: max_tokens=%d field=%q", found.MaxTokens, found.MaxTokensField)
}
// Update with a blank key takes the "keep existing key" UPDATE branch, which
// has its own column list and is the easiest one to forget.
p.APIKey = ""
p.MaxTokens = 4096
p.MaxTokensField = ""
if _, err := d.SaveProfile(p); err != nil {
t.Fatal(err)
}
after, err := d.ProfileByID(id)
if err != nil || after == nil {
t.Fatalf("reload: %v", err)
}
if after.MaxTokens != 4096 || after.MaxTokensField != "" {
t.Fatalf("after update: max_tokens=%d field=%q", after.MaxTokens, after.MaxTokensField)
}
if after.APIKey != "k" {
t.Fatalf("blank key on update must keep the stored one, got %q", after.APIKey)
}
}
// A profile saved without touching the new fields must read back as "no cap,
// classic field name" — the pre-feature behaviour old rows also get.
func TestProfileMaxTokensDefaults(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
// Close via Cleanup, registered FIRST so it runs LAST: cleanups are LIFO, and a
// plain `defer d.Close()` would fire before them — the row-deleting cleanups
// would then run against a closed pool and silently leave test rows behind.
t.Cleanup(func() { d.Close() })
id, err := d.SaveProfile(&LLMProfile{Name: "t-maxtok-default", Format: "openai", Model: "m", APIKey: "k"})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { d.Exec(`DELETE FROM llm_profiles WHERE id=$1`, id) })
p, err := d.ProfileByID(id)
if err != nil || p == nil {
t.Fatalf("ProfileByID: %v", err)
}
if p.MaxTokens != 0 || p.MaxTokensField != "" {
t.Fatalf("defaults: max_tokens=%d field=%q, want 0 and \"\"", p.MaxTokens, p.MaxTokensField)
}
}
+196
View File
@@ -0,0 +1,196 @@
package db
import (
"context"
"encoding/json"
"errors"
"testing"
)
// TestPoolProfilesOrder pins the failover chain query: keyless profiles can't
// serve a request and excluded ones aren't fallback targets, so neither belongs
// in the chain; the rest come back by priority, highest first.
// Deliberately does NOT touch is_default — flipping the active profile would be a
// side effect on the shared dev database.
func TestPoolProfilesOrder(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
mk := func(name string, priority int, exclude bool, key string) int64 {
id, err := d.SaveProfile(&LLMProfile{
Name: name, Format: "openai", Model: "m", APIKey: key,
Priority: priority, PoolExclude: exclude,
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { d.Exec(`DELETE FROM llm_profiles WHERE id=$1`, id) })
return id
}
lo := mk("t-pool-lo", 1, false, "k1")
hi := mk("t-pool-hi", 9, false, "k2")
mk("t-pool-excluded", 99, true, "k3") // excluded despite the top priority
mk("t-pool-nokey", 50, false, "") // no key → cannot serve anything
chain, err := d.PoolProfiles()
if err != nil {
t.Fatal(err)
}
var got []int64
for _, p := range chain {
switch p.Name {
case "t-pool-lo", "t-pool-hi":
got = append(got, p.ID)
case "t-pool-excluded":
t.Fatal("pool_exclude profile entered the failover chain")
case "t-pool-nokey":
t.Fatal("keyless profile entered the failover chain")
}
}
if len(got) != 2 || got[0] != hi || got[1] != lo {
t.Fatalf("chain order = %v, want [hi=%d lo=%d]", got, hi, lo)
}
}
func TestDeleteProfileContextHonorsCancellation(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
var d DB
if err := d.DeleteProfileContext(ctx, 1); !errors.Is(err, context.Canceled) {
t.Fatalf("DeleteProfileContext error=%v, want context cancellation", err)
}
}
func TestConfigStores(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
// LLM profile: save, active, key never serialized
pid, err := d.SaveProfile(&LLMProfile{Name: "t-default", Format: "openai", Model: "gpt-x", APIKey: "secret123"})
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM llm_profiles WHERE id=$1`, pid)
if err := d.SetActiveProfile(pid); err != nil {
t.Fatal(err)
}
if err := d.DeleteProfile(pid); !errors.Is(err, ErrActiveLLMProfileDelete) {
t.Fatalf("deleting active profile error=%v, want %v", err, ErrActiveLLMProfileDelete)
}
act, err := d.ActiveProfile()
if err != nil || act == nil || act.APIKey != "secret123" {
t.Fatalf("active profile/key: %+v err=%v", act, err)
}
// list must hide the key, expose hint
list, _ := d.ListProfiles()
for _, p := range list {
if p.ID == pid {
b, _ := json.Marshal(p)
if string(b) == "" || contains(string(b), "secret123") {
t.Fatalf("api key leaked in list json: %s", b)
}
if p.APIKeyHint != "…t123" {
t.Fatalf("hint want …t123, got %q", p.APIKeyHint)
}
}
}
// agents seeded; prompt versioning
ag, err := d.GetAgentByKey("planner")
if err != nil || ag == nil {
t.Fatalf("planner agent: %v", err)
}
v1, err := d.SavePrompt(ag.ID, "你是规划者 {{.Goal}}", "init", "test")
if err != nil {
t.Fatal(err)
}
v2, _ := d.SavePrompt(ag.ID, "你是规划者 v2 {{.Goal}} {{.Scope}}", "edit", "test")
if v2 != v1+1 {
t.Fatalf("version should increment: %d -> %d", v1, v2)
}
cur, _ := d.CurrentPrompt(ag.ID)
if cur != "你是规划者 v2 {{.Goal}} {{.Scope}}" {
t.Fatalf("current prompt wrong: %q", cur)
}
vers, _ := d.ListPromptVersions(ag.ID)
if len(vers) < 2 {
t.Fatalf("want >=2 versions, got %d", len(vers))
}
pv, _ := d.PromptVars(ag.ID)
// planner must have at least the seeded catalog vars (Goal, AssetSummary)
if len(pv) < 2 {
t.Fatalf("planner catalog want >=2 vars, got %d", len(pv))
}
d.Exec(`DELETE FROM agent_prompts WHERE agent_id=$1`, ag.ID)
d.Exec(`UPDATE agents SET current_prompt_id=NULL WHERE id=$1`, ag.ID)
// mcp + skill + visibility (bidirectional via one join)
// Clean up any leftover MCP from prior runs to keep this test idempotent.
d.Exec(`DELETE FROM mcp_servers WHERE name = 't-gh'`)
mid, err := d.SaveMCP(&MCPServer{Name: "t-gh", Transport: "stdio", Command: "npx", Args: json.RawMessage(`["server-github"]`), Env: json.RawMessage(`{"GITHUB_TOKEN":"x"}`), Enabled: true})
if err != nil {
t.Fatal(err)
}
// agent-side write. MCP is id-keyed (generic visibility join); skills are now
// filesystem-based, so their visibility is keyed by skill (directory) name in a
// dedicated table. Clear any pre-existing MCP visibility rows so the assertions
// below isolate on exactly what this test sets.
d.Exec(`DELETE FROM agent_visibility WHERE agent_id = $1 AND resource_kind = 'mcp'`, ag.ID)
if err := d.ToggleVisibility(ag.ID, "mcp", mid, true); err != nil {
t.Fatal(err)
}
if err := d.SetAgentSkillVisibility(ag.ID, nil); err != nil {
t.Fatal(err)
}
if err := d.ToggleSkillVisibility(ag.ID, "t-skill", true); err != nil {
t.Fatal(err)
}
// agent-side read
vm, _ := d.AgentVisible(ag.ID, "mcp")
if len(vm) != 1 || vm[0] != mid {
t.Fatalf("agent visible mcp: %+v", vm)
}
// resource-side read (same join row) → bidirectional
ra, _ := d.ResourceAgents("mcp", mid)
if len(ra) != 1 || ra[0] != ag.ID {
t.Fatalf("resource agents: %+v", ra)
}
// toggle off
d.ToggleVisibility(ag.ID, "mcp", mid, false)
vm2, _ := d.AgentVisible(ag.ID, "mcp")
if len(vm2) != 0 {
t.Fatalf("after toggle off: %+v", vm2)
}
// skill visibility is name-keyed: verify the read, then deleting the skill's
// visibility rows (called when a skill is removed) clears it.
names, _ := d.AgentSkillNames(ag.ID)
if len(names) != 1 || names[0] != "t-skill" {
t.Fatalf("agent visible skills: %+v", names)
}
if err := d.DeleteSkillVisibility("t-skill"); err != nil {
t.Fatal(err)
}
names2, _ := d.AgentSkillNames(ag.ID)
if len(names2) != 0 {
t.Fatalf("skill visibility should be cleared on delete: %+v", names2)
}
d.DeleteMCP(mid)
}
func contains(s, sub string) bool {
return len(s) >= len(sub) && (indexOf(s, sub) >= 0)
}
func indexOf(s, sub string) int {
for i := 0; i+len(sub) <= len(s); i++ {
if s[i:i+len(sub)] == sub {
return i
}
}
return -1
}
+40
View File
@@ -0,0 +1,40 @@
package db
// Exploration node kinds (exploration_nodes.kind).
const (
KindBegin = "begin" // DEPRECATED: legacy task root; new tasks seed an origin fact (KindFact + StateOrigin) instead
KindGoal = "goal" // a task objective
KindIntent = "intent" // an exploration direction (planner-generated)
KindFact = "fact" // a worker's exploration result/conclusion (incl. negative results), tied to its intent
KindFinding = "finding" // a confirmed vulnerability (report_finding), distinct from a fact
KindHint = "hint"
KindDigest = "digest" // a compressed fold of cold intents/facts (cold-digest-spec §1); lossless — members kept, restorable by id
)
// Digest node states (kind='digest'). A digest is 'active' while it renders in
// graph_overview; major compaction retires a merged-away segment to 'superseded'
// (its covers edges repointed to the new digest) — cold-digest-spec §5.1.
const (
StateDigestActive = "active"
StateDigestSuperseded = "superseded"
)
// StateOrigin marks the task root fact (KindFact) seeded at task creation — the
// exploration graph's origin. Every intent traces back to it, so "an intent must
// connect to a fact" holds uniformly from the very first intent. Worker-produced
// facts use state 'confirmed', so this never collides.
const StateOrigin = "origin"
// StateIntentDeleted marks an intent the user假删除(soft delete): it drops out of
// the frontier and graph_overview like other terminal states, but keeps its node
// and full lineage. The delete reason lives in exploration_nodes.delete_reason.
const StateIntentDeleted = "deleted"
// Exploration edge relations (exploration_edges.rel).
const (
RelSpawns = "spawns"
RelDerivedFrom = "derived_from"
RelYields = "yields"
RelProves = "proves"
RelCovers = "covers" // digest --covers--> member (cold-digest-spec §1); source of truth for "which digest folds node X"
)
+85
View File
@@ -0,0 +1,85 @@
package db
import (
"fmt"
"time"
)
// Constraint is one operator-authored operation constraint for a task: kind=allow
// (permitted operations) or kind=deny (forbidden operations), free-text. Stored in
// task_constraints, keyed by exploration_id (cascades with the exploration).
type Constraint struct {
ID int64 `json:"id"`
Kind string `json:"kind"` // allow | deny
Text string `json:"text"`
Origin string `json:"origin,omitempty"` // goals | human | system
CreatedAt time.Time `json:"created_at"`
}
// ListConstraints returns this exploration's constraints, allow before deny, oldest
// first within each group (stable render order for the prompt block + UI).
func (s *ExplorationStore) ListConstraints() ([]Constraint, error) {
rows, err := s.db.Query(`
SELECT id, kind, text, COALESCE(origin,''), created_at
FROM task_constraints WHERE exploration_id=$1
ORDER BY (kind='deny'), id`, s.expID)
if err != nil {
return nil, err
}
defer rows.Close()
var out []Constraint
for rows.Next() {
var c Constraint
if err := rows.Scan(&c.ID, &c.Kind, &c.Text, &c.Origin, &c.CreatedAt); err != nil {
return nil, err
}
out = append(out, c)
}
return out, rows.Err()
}
// AddConstraint inserts one constraint (kind must be allow|deny) and returns its id.
func (s *ExplorationStore) AddConstraint(kind, text, origin string) (int64, error) {
if kind != "allow" && kind != "deny" {
return 0, fmt.Errorf("kind 必须是 allow 或 deny")
}
if origin == "" {
origin = "system"
}
var id int64
err := s.db.QueryRow(`
INSERT INTO task_constraints(exploration_id, kind, text, origin)
VALUES ($1, $2, $3, $4) RETURNING id`, s.expID, kind, text, origin).Scan(&id)
return id, err
}
// UpdateConstraint rewrites a constraint's kind + text; scoped to this exploration.
// Returns an error if no such constraint exists.
func (s *ExplorationStore) UpdateConstraint(id int64, kind, text string) error {
if kind != "allow" && kind != "deny" {
return fmt.Errorf("kind 必须是 allow 或 deny")
}
res, err := s.db.Exec(`
UPDATE task_constraints SET kind=$1, text=$2, updated_at=now()
WHERE id=$3 AND exploration_id=$4`, kind, text, id, s.expID)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return fmt.Errorf("约束不存在")
}
return nil
}
// DeleteConstraint removes a constraint; scoped to this exploration. Returns an
// error if no such constraint exists.
func (s *ExplorationStore) DeleteConstraint(id int64) error {
res, err := s.db.Exec(`DELETE FROM task_constraints WHERE id=$1 AND exploration_id=$2`, id, s.expID)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return fmt.Errorf("约束不存在")
}
return nil
}
+327
View File
@@ -0,0 +1,327 @@
package db
import (
"database/sql"
"time"
)
// ConvTokenSummary is one conversation's token total (sum of its kind='result'
// rows) with its profile + created_at, used to merge conversation usage into the
// dashboard's per-profile / daily token stats (which otherwise cover only tasks).
type ConvTokenSummary struct {
LLMProfileID *int64 `json:"llm_profile_id"`
CreatedAt string `json:"created_at"`
InputTokens int `json:"input_tokens"`
OutputTokens int `json:"output_tokens"`
CacheReadTokens int `json:"cache_read_tokens"`
CacheWriteTokens int `json:"cache_write_tokens"`
}
// ConversationTokenSummaries returns one row per conversation with its summed
// result-row token usage (0 for conversations with no completed run yet).
func (d *DB) ConversationTokenSummaries() ([]ConvTokenSummary, error) {
rows, err := d.Query(`
SELECT c.llm_profile_id, c.created_at::text,
COALESCE(sum(ca.input_tokens),0), COALESCE(sum(ca.output_tokens),0),
COALESCE(sum(ca.cache_read_tokens),0), COALESCE(sum(ca.cache_write_tokens),0)
FROM conversations c
LEFT JOIN conversation_activities ca ON ca.conversation_id = c.id AND ca.kind = 'result'
GROUP BY c.id`)
if err != nil {
return nil, err
}
defer rows.Close()
out := []ConvTokenSummary{}
for rows.Next() {
var s ConvTokenSummary
var pid sql.NullInt64
if err := rows.Scan(&pid, &s.CreatedAt, &s.InputTokens, &s.OutputTokens, &s.CacheReadTokens, &s.CacheWriteTokens); err != nil {
return nil, err
}
if pid.Valid {
v := pid.Int64
s.LLMProfileID = &v
}
out = append(out, s)
}
return out, rows.Err()
}
// Conversation is one ChatGPT-style chat thread bound to an agent key. It lives
// independent of the pentest exploration graph — see schema.sql §I.
type Conversation struct {
ID int64 `json:"id"`
AgentKey string `json:"agent_key"`
Title string `json:"title"`
LLMProfileID *int64 `json:"llm_profile_id,omitempty"`
Pinned bool `json:"pinned"`
PinnedAt *time.Time `json:"pinned_at,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// ConversationPatch updates only the fields whose pointers are non-nil.
type ConversationPatch struct {
Title *string
Pinned *bool
}
const convCols = `id, agent_key, title, llm_profile_id, pinned_at, created_at, updated_at`
func scanConv(row interface{ Scan(...any) error }) (Conversation, error) {
var c Conversation
var pinnedAt sql.NullTime
err := row.Scan(&c.ID, &c.AgentKey, &c.Title, &c.LLMProfileID, &pinnedAt, &c.CreatedAt, &c.UpdatedAt)
if pinnedAt.Valid {
c.Pinned = true
c.PinnedAt = &pinnedAt.Time
}
return c, err
}
// CreateConversation opens a new chat thread for agentKey with an initial title.
// llmProfileID may be nil to use the globally active profile.
func (d *DB) CreateConversation(agentKey, title string, llmProfileID *int64) (*Conversation, error) {
tx, err := d.Begin()
if err != nil {
return nil, err
}
defer tx.Rollback()
// Insert the child without a profile first. The new row is exclusively owned
// by this transaction before it takes a profile lock, matching DeleteProfile's
// child-row -> profile-row protocol.
c, err := scanConv(tx.QueryRow(`
INSERT INTO conversations(agent_key, title, llm_profile_id) VALUES ($1, $2, NULL)
RETURNING `+convCols, agentKey, title))
if err != nil {
return nil, err
}
if err := lockLLMProfileForReference(tx, llmProfileID); err != nil {
return nil, err
}
if llmProfileID != nil {
c, err = scanConv(tx.QueryRow(`UPDATE conversations SET llm_profile_id=$2
WHERE id=$1 RETURNING `+convCols, c.ID, llmProfileID))
if err != nil {
return nil, err
}
}
if err := tx.Commit(); err != nil {
return nil, err
}
return &c, nil
}
// UpdateConversationProfile sets (or clears) the LLM profile override for a conversation.
func (d *DB) UpdateConversationProfile(id int64, llmProfileID *int64) error {
tx, err := d.Begin()
if err != nil {
return err
}
defer tx.Rollback()
var lockedID int64
if err := tx.QueryRow(`SELECT id FROM conversations WHERE id=$1 FOR UPDATE`, id).Scan(&lockedID); err != nil {
if err == sql.ErrNoRows {
// Preserve the previous UPDATE semantics: an unknown id is a no-op.
return tx.Commit()
}
return err
}
if err := lockLLMProfileForReference(tx, llmProfileID); err != nil {
return err
}
if _, err := tx.Exec(`UPDATE conversations SET llm_profile_id=$2 WHERE id=$1`, id, llmProfileID); err != nil {
return err
}
return tx.Commit()
}
// ListConversations returns all threads, most-recently-updated first.
func (d *DB) ListConversations() ([]*Conversation, error) {
rows, err := d.Query(`SELECT ` + convCols + ` FROM conversations
ORDER BY (pinned_at IS NOT NULL) DESC, pinned_at DESC NULLS LAST, updated_at DESC, id DESC`)
if err != nil {
return nil, err
}
defer rows.Close()
out := []*Conversation{}
for rows.Next() {
c, err := scanConv(rows)
if err != nil {
return nil, err
}
out = append(out, &c)
}
return out, rows.Err()
}
// GetConversation returns one thread (nil, nil if absent).
func (d *DB) GetConversation(id int64) (*Conversation, error) {
c, err := scanConv(d.QueryRow(`SELECT `+convCols+` FROM conversations WHERE id=$1`, id))
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
return &c, nil
}
// UpdateConversation applies a partial title/pin mutation and returns the updated
// row. Pinning an already-pinned conversation preserves its original pin order.
func (d *DB) UpdateConversation(id int64, patch ConversationPatch) (*Conversation, error) {
c, err := scanConv(d.QueryRow(`UPDATE conversations SET
title = CASE WHEN $2::boolean THEN $3 ELSE title END,
pinned_at = CASE
WHEN $4::boolean IS NULL THEN pinned_at
WHEN $4::boolean THEN COALESCE(pinned_at, now())
ELSE NULL
END
WHERE id=$1
RETURNING `+convCols, id, patch.Title != nil, patch.Title, patch.Pinned))
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
return &c, nil
}
// RenameConversation sets a thread's title. Kept for automatic first-message
// titles and compatibility with existing callers.
func (d *DB) RenameConversation(id int64, title string) error {
_, err := d.UpdateConversation(id, ConversationPatch{Title: &title})
return err
}
// TouchConversation bumps updated_at so the thread floats to the top of the list.
func (d *DB) TouchConversation(id int64) error {
_, err := d.Exec(`UPDATE conversations SET updated_at=now() WHERE id=$1`, id)
return err
}
// DeleteConversation removes a thread; its activities cascade via FK.
func (d *DB) DeleteConversation(id int64) error {
_, err := d.Exec(`DELETE FROM conversations WHERE id=$1`, id)
return err
}
// DeleteConversations removes existing threads in one statement and returns the
// ids that were actually present. Child activities and trigger runs cascade.
func (d *DB) DeleteConversations(ids []int64) ([]int64, error) {
if len(ids) == 0 {
return []int64{}, nil
}
rows, err := d.Query(`DELETE FROM conversations WHERE id=ANY($1::bigint[]) RETURNING id`, ids)
if err != nil {
return nil, err
}
defer rows.Close()
deleted := make([]int64, 0, len(ids))
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
return nil, err
}
deleted = append(deleted, id)
}
return deleted, rows.Err()
}
// AppendConvActivity records one step of a conversation (human message or an agent
// execution step) and returns its id. Mirrors ExplorationStore.AppendActivity but
// keyed by conversation_id. Reuses the Activity struct (NodeID is ignored here).
func (d *DB) AppendConvActivity(convID int64, a Activity) (int64, error) {
var id int64
err := d.QueryRow(`
INSERT INTO conversation_activities(conversation_id, worker, kind, tool, tool_use_id, is_error, summary, detail, input_tokens, output_tokens, cache_read_tokens, cache_write_tokens)
VALUES ($1,NULLIF($2,''),NULLIF($3,''),NULLIF($4,''),NULLIF($5,''),$6,NULLIF($7,''),NULLIF($8,''),$9,$10,$11,$12)
RETURNING id`, convID, utf8Clean(a.Worker), utf8Clean(a.Kind), utf8Clean(a.Tool), utf8Clean(a.ToolUseID), a.IsError,
utf8Clean(a.Summary), utf8Clean(a.Detail), a.InputTokens, a.OutputTokens, a.CacheReadTokens, a.CacheWriteTokens).Scan(&id)
return id, err
}
// ConvActivityList returns a conversation's steps after sinceID (exclusive) with
// the summary-only column set (detail is lazy-loaded via ConvActivityDetail).
func (d *DB) ConvActivityList(convID, sinceID int64, limit int) ([]Activity, int64, error) {
if limit <= 0 {
limit = 500
}
const cols = `id, COALESCE(worker,''), COALESCE(kind,''), COALESCE(tool,''), COALESCE(tool_use_id,''), is_error, COALESCE(summary,''), created_at, input_tokens, output_tokens, cache_read_tokens, cache_write_tokens`
rows, err := d.Query(`SELECT `+cols+`
FROM conversation_activities WHERE conversation_id=$1 AND id>$2 ORDER BY id LIMIT $3`, convID, sinceID, limit)
if err != nil {
return nil, sinceID, err
}
defer rows.Close()
out := []Activity{}
cursor := sinceID
for rows.Next() {
var a Activity
if err := rows.Scan(&a.ID, &a.Worker, &a.Kind, &a.Tool, &a.ToolUseID, &a.IsError, &a.Summary, &a.CreatedAt,
&a.InputTokens, &a.OutputTokens, &a.CacheReadTokens, &a.CacheWriteTokens); err != nil {
return nil, sinceID, err
}
if a.ID > cursor {
cursor = a.ID
}
out = append(out, a)
}
return out, cursor, rows.Err()
}
// ConvActivityPage returns one page for reverse (newest-first) pagination: up to
// `limit` steps ending before id `before` (exclusive; before<=0 = the latest
// page), returned in ASCENDING id order. hasMore reports whether still-older steps
// exist before the returned window, so the client can stop loading earlier history
// on scroll-up. Summary-only columns (detail is lazy-loaded via ConvActivityDetail).
func (d *DB) ConvActivityPage(convID, before int64, limit int) ([]Activity, bool, error) {
if limit <= 0 {
limit = 200
}
const cols = `id, COALESCE(worker,''), COALESCE(kind,''), COALESCE(tool,''), COALESCE(tool_use_id,''), is_error, COALESCE(summary,''), created_at, input_tokens, output_tokens, cache_read_tokens, cache_write_tokens`
// fetch one extra row to detect whether older history remains before this window.
rows, err := d.Query(`SELECT `+cols+`
FROM conversation_activities
WHERE conversation_id=$1 AND ($2 <= 0 OR id < $2)
ORDER BY id DESC LIMIT $3`, convID, before, limit+1)
if err != nil {
return nil, false, err
}
defer rows.Close()
desc := []Activity{}
for rows.Next() {
var a Activity
if err := rows.Scan(&a.ID, &a.Worker, &a.Kind, &a.Tool, &a.ToolUseID, &a.IsError, &a.Summary, &a.CreatedAt,
&a.InputTokens, &a.OutputTokens, &a.CacheReadTokens, &a.CacheWriteTokens); err != nil {
return nil, false, err
}
desc = append(desc, a)
}
if err := rows.Err(); err != nil {
return nil, false, err
}
hasMore := len(desc) > limit
if hasMore {
desc = desc[:limit]
}
// reverse the newest-first window into ascending id order for display.
out := make([]Activity, len(desc))
for i, a := range desc {
out[len(desc)-1-i] = a
}
return out, hasMore, nil
}
// ConvActivityDetail lazily returns the full detail blob for one step.
func (d *DB) ConvActivityDetail(convID, id int64) (string, error) {
var s sql.NullString
err := d.QueryRow(`SELECT detail FROM conversation_activities WHERE id=$1 AND conversation_id=$2`, id, convID).Scan(&s)
if err == sql.ErrNoRows {
return "", nil
}
return s.String, err
}
+69
View File
@@ -0,0 +1,69 @@
package db
import (
"fmt"
"testing"
"time"
)
func TestConversationPinOrderingAndPatch(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
suffix := time.Now().UnixNano()
first, err := d.CreateConversation("mainagent", fmt.Sprintf("pin-first-%d", suffix), nil)
if err != nil {
t.Fatal(err)
}
second, err := d.CreateConversation("mainagent", fmt.Sprintf("pin-second-%d", suffix), nil)
if err != nil {
_ = d.DeleteConversation(first.ID)
t.Fatal(err)
}
t.Cleanup(func() {
_ = d.DeleteConversation(first.ID)
_ = d.DeleteConversation(second.ID)
})
pinned := true
first, err = d.UpdateConversation(first.ID, ConversationPatch{Pinned: &pinned})
if err != nil || first == nil || !first.Pinned || first.PinnedAt == nil {
t.Fatalf("pin first = %+v, %v", first, err)
}
firstPinnedAt := *first.PinnedAt
time.Sleep(5 * time.Millisecond)
second, err = d.UpdateConversation(second.ID, ConversationPatch{Pinned: &pinned})
if err != nil || second == nil || !second.Pinned {
t.Fatalf("pin second = %+v, %v", second, err)
}
newTitle := "renamed while pinned"
first, err = d.UpdateConversation(first.ID, ConversationPatch{Title: &newTitle, Pinned: &pinned})
if err != nil || first == nil {
t.Fatal(err)
}
if first.Title != newTitle || first.PinnedAt == nil || !first.PinnedAt.Equal(firstPinnedAt) {
t.Fatalf("repeat pin should preserve pin time: %+v (want %v)", first, firstPinnedAt)
}
conversations, err := d.ListConversations()
if err != nil {
t.Fatal(err)
}
positions := map[int64]int{}
for index, conversation := range conversations {
positions[conversation.ID] = index
}
if positions[second.ID] >= positions[first.ID] {
t.Fatalf("newer pin must sort first: second=%d first=%d", positions[second.ID], positions[first.ID])
}
pinned = false
first, err = d.UpdateConversation(first.ID, ConversationPatch{Pinned: &pinned})
if err != nil || first == nil || first.Pinned || first.PinnedAt != nil {
t.Fatalf("unpin first = %+v, %v", first, err)
}
}
+82
View File
@@ -0,0 +1,82 @@
package db
import (
"encoding/json"
"testing"
)
// TestCustomToolCRUD exercises create/list/update/delete of a user-defined tool
// (system=false) with kind/exec/deferred. Skips when no Postgres is configured.
func TestCustomToolCRUD(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
key := "ct_test_tool"
_ = d.DeleteCustomTool(key) // clean slate
in := &Tool{
Key: key,
Description: "test",
Schema: json.RawMessage(`{"type":"object","properties":{"x":{"type":"string"}}}`),
Agents: []string{"worker"},
Enabled: true,
Kind: "script",
Exec: json.RawMessage(`{"code":"print(1)"}`),
Deferred: true,
}
if err := d.CreateCustomTool(in); err != nil {
t.Fatalf("create: %v", err)
}
t.Cleanup(func() { _ = d.DeleteCustomTool(key) })
got, err := d.GetTool(key)
if err != nil || got == nil {
t.Fatalf("get after create: %v (nil=%v)", err, got == nil)
}
if got.System || got.Kind != "script" || !got.Deferred || len(got.Agents) != 1 {
t.Fatalf("unexpected row: %+v", got)
}
// custom tools appear in ListCustomTools, not among system-only.
customs, err := d.ListCustomTools()
if err != nil {
t.Fatal(err)
}
found := false
for _, c := range customs {
if c.Key == key {
found = true
if c.System {
t.Fatal("custom tool should have system=false")
}
}
}
if !found {
t.Fatal("created tool not in ListCustomTools")
}
// update: change kind + disable + rebind.
in.Kind = "command"
in.Exec = json.RawMessage(`{"command":"echo hi"}`)
in.Enabled = false
in.Agents = []string{"worker", "planner"}
in.Deferred = false
if err := d.UpdateCustomTool(in); err != nil {
t.Fatalf("update: %v", err)
}
got2, _ := d.GetTool(key)
if got2.Kind != "command" || got2.Enabled || got2.Deferred || len(got2.Agents) != 2 {
t.Fatalf("update not applied: %+v", got2)
}
// delete only removes non-system rows.
if err := d.DeleteCustomTool(key); err != nil {
t.Fatalf("delete: %v", err)
}
if g, _ := d.GetTool(key); g != nil {
t.Fatal("tool still present after delete")
}
}
+602
View File
@@ -0,0 +1,602 @@
// Package db is the PostgreSQL data source for ARTEX (取代旧 graph 单文件 SQLite)。
// 它打开连接、应用 schema、并 seed 内置 agent 与变量目录。
package db
import (
"context"
"database/sql"
_ "embed"
"errors"
"fmt"
"net/url"
"strings"
"time"
"github.com/Autumn-27/artex/config"
"github.com/jackc/pgx/v5/pgconn"
_ "github.com/jackc/pgx/v5/stdlib" // pgx database/sql driver ("pgx")
)
//go:embed schema.sql
var schemaSQL string
const schemaMigrationLockKey int64 = 7337741001
var schemaDeadlockRetryDelays = [...]time.Duration{
100 * time.Millisecond,
250 * time.Millisecond,
500 * time.Millisecond,
time.Second,
}
type schemaExecer interface {
ExecContext(context.Context, string, ...any) (sql.Result, error)
}
func isPostgresDeadlock(err error) bool {
var pgErr *pgconn.PgError
return errors.As(err, &pgErr) && pgErr.Code == "40P01"
}
func applySchemaWithRetry(ctx context.Context, execer schemaExecer, sleep func(time.Duration)) error {
for attempt := 0; ; attempt++ {
if _, err := execer.ExecContext(ctx, schemaSQL); err != nil {
if !isPostgresDeadlock(err) || attempt >= len(schemaDeadlockRetryDelays) {
return err
}
sleep(schemaDeadlockRetryDelays[attempt])
continue
}
return nil
}
}
// withSchemaMigrationLock pins the session-level lock to one checked-out
// connection. Running pg_advisory_lock through *sql.DB is incorrect because a
// later schema or unlock call may use a different pooled PostgreSQL session.
func withSchemaMigrationLock(ctx context.Context, sqlDB *sql.DB, action func(*sql.Conn) error) (err error) {
conn, err := sqlDB.Conn(ctx)
if err != nil {
return err
}
defer conn.Close()
if _, err := conn.ExecContext(ctx, `SELECT pg_advisory_lock($1)`, schemaMigrationLockKey); err != nil {
return fmt.Errorf("advisory lock: %w", err)
}
defer func() {
if _, unlockErr := conn.ExecContext(context.Background(), `SELECT pg_advisory_unlock($1)`, schemaMigrationLockKey); unlockErr != nil && err == nil {
err = fmt.Errorf("advisory unlock: %w", unlockErr)
}
}()
return action(conn)
}
// coordinateWithSchemaMigration makes long, multi-table archive transactions
// mutually exclusive with startup DDL while allowing ordinary runtime queries
// to continue normally.
func coordinateWithSchemaMigration(tx *sql.Tx) error {
if _, err := tx.Exec(`SELECT pg_advisory_xact_lock($1)`, schemaMigrationLockKey); err != nil {
return fmt.Errorf("coordinate with schema migration: %w", err)
}
return nil
}
// DSN resolves the PostgreSQL connection string and reports where it came from.
// Precedence: env ARTEX_PG_DSN > config file (config.json). There is no
// built-in default — it errors if neither source is configured.
func DSN() (dsn, source string, err error) {
return config.PostgresDSN()
}
// DB wraps the shared *sql.DB. PG handles its own connection pool + concurrency
// (MVCC), so unlike the old SQLite store there is no process-wide write mutex.
type DB struct{ *sql.DB }
// ensureDatabase connects to the postgres system database and creates the target
// database if it does not exist. dsn must be a postgres:// URL.
func ensureDatabase(dsn string) error {
u, err := url.Parse(dsn)
if err != nil {
return nil // unparseable DSN — let the normal Open fail with a clear error
}
dbName := strings.TrimPrefix(u.Path, "/")
if dbName == "" || dbName == "postgres" {
return nil
}
// connect to the postgres maintenance database instead
adminDSN := *u
adminDSN.Path = "/postgres"
admin, err := sql.Open("pgx", adminDSN.String())
if err != nil {
return nil // best-effort; let Open surface the real error
}
defer admin.Close()
if err := admin.Ping(); err != nil {
return nil
}
var exists bool
_ = admin.QueryRow(`SELECT true FROM pg_database WHERE datname=$1`, dbName).Scan(&exists)
if !exists {
if _, err := admin.Exec(`CREATE DATABASE "` + dbName + `"`); err != nil {
return fmt.Errorf("create database %q: %w", dbName, err)
}
}
return nil
}
// Open connects, applies the schema (idempotent), and seeds builtin rows.
func Open(dsn string) (*DB, error) {
if err := ensureDatabase(dsn); err != nil {
return nil, err
}
sqlDB, err := sql.Open("pgx", dsn)
if err != nil {
return nil, err
}
if err := sqlDB.Ping(); err != nil {
sqlDB.Close()
return nil, fmt.Errorf("ping postgres (%s): %w", config.Redact(dsn), err)
}
d := &DB{sqlDB}
// pgx runs multi-statement Exec via the simple protocol when there are no args.
// Keep the dedicated lock connection checked out until both DDL and seeding
// finish so concurrent application instances cannot initialize out of order.
err = withSchemaMigrationLock(context.Background(), sqlDB, func(conn *sql.Conn) error {
if err := applySchemaWithRetry(context.Background(), conn, time.Sleep); err != nil {
return fmt.Errorf("apply schema: %w", err)
}
if err := d.seedBuiltins(); err != nil {
return fmt.Errorf("seed builtins: %w", err)
}
return nil
})
if err != nil {
sqlDB.Close()
return nil, err
}
return d, nil
}
// builtinAgent describes one of the fixed agents and its prompt-variable catalog.
type builtinAgent struct {
key, name, role, desc string
vars []promptVar
interactiveShell bool // 建行时的默认交互式 shell 开关;ON CONFLICT 不覆盖用户后续手动开关
runSeconds *int // 建行时的单次 run 墙钟上限(秒);nil=用种子默认(1200),0=不限时
}
type promptVar struct{ name, desc, example, source string }
// intp 返回 v 的指针,用于给 builtinAgent 可选字段(如 runSeconds)显式取值。
func intp(v int) *int { return &v }
// builtinAgents mirrors docs §5(a). 内置工具不入库;这里只 seed agent + 变量目录。
// 注:planner/worker/mainagent/auto 的交互式 shell 默认由下方 interactive_shell_default_v1
// 块统一置 true(尊重后续 toggle);这里的 interactiveShell 只给需要「建行即默认开」的新 agent。
var builtinAgents = []builtinAgent{
{"goals", "目标拆解", "goals", "把渗透任务目标拆解成若干独立、可验证的子目标。", []promptVar{
{"EngagementDescription", "任务描述(测试对象/背景)", "测试 example.com 站点", "exploration"},
// Now 是全局 runtime 变量(见 server.globalPromptVars),不再在各 agent 目录里
// 重复定义,否则 withGlobalVars 追加时会与全局项撞名。
}, false, nil},
{"planner", "规划", "planner", "读取态势、判定目标,只在确有未覆盖的新方向时补充探索意图(每任务一个规划循环)。", []promptVar{
{"Goal", "任务总目标", "拿下 example.com 的管理员权限", "exploration"},
{"AssetSummary", "资产计数/类型分布摘要(可选)", "domain:3 ip:5 site:2", "distilled"},
}, false, nil},
{"mainagent", "主", "main", "人机接口:观察进展,把人的意图落成 hint 或高优先级意图。", []promptVar{
{"Goal", "当前任务目标", "拿下 example.com 的管理员权限", "exploration"},
{"AssetSummary", "开局态势摘要(可选)", "domain:3 ip:5", "distilled"},
{"FindingsSummary", "已确认漏洞摘要(可选)", "high:1 medium:2", "distilled"},
}, false, nil},
{"worker", "执行", "worker", "领取一条意图执行,把发现的事实/漏洞写回知识图谱后停止。", []promptVar{
{"ProxyAddr", "记录代理地址(驱动 if 双文案)", "127.0.0.1:8080", "runtime"},
{"WorkerName", "worker 自我标识(可选)", "worker-1", "runtime"},
}, false, nil},
// Auto:内置「平台操作」agent。不参与渗透编排循环,经对话页驱动,用工具操作平台。
{"auto", "Auto", "assistant", "平台操作助手:用工具管理任务(建/看/暂停/给提示)与资产,并可创建/修改 skill、自定义工具、MCP。", nil, false, nil},
// 渗透测试:内置「独立渗透」agent。经对话页驱动,一人从侦察到收尾走完整条渗透链,自己规划自己执行自己验证。默认开启交互式 shell。
{"pentest", "渗透测试", "assistant", "独立渗透 agent:一人从侦察→找攻击面→深入利用→验证→收尾走完整条链,自己规划、自己执行、自己对抗式验证。", nil, true, intp(0)},
}
// seedBuiltins inserts the fixed built-in agents and their variable catalog (idempotent).
func (d *DB) seedBuiltins() error {
for _, a := range builtinAgents {
var agentID int64
err := d.QueryRow(`
INSERT INTO agents(key, name, description, role, builtin, enabled, interactive_shell, run_seconds)
VALUES ($1, $2, NULLIF($3,''), $4, true, true, $5, COALESCE($6, 1200))
ON CONFLICT (key) DO UPDATE SET name = EXCLUDED.name, description = EXCLUDED.description
RETURNING id`, a.key, a.name, a.desc, a.role, a.interactiveShell, a.runSeconds).Scan(&agentID)
if err != nil {
return fmt.Errorf("agent %s: %w", a.key, err)
}
for _, v := range a.vars {
if _, err := d.Exec(`
INSERT INTO agent_prompt_vars(agent_id, var_name, description, example, source)
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (agent_id, var_name) DO UPDATE
SET description = EXCLUDED.description, example = EXCLUDED.example, source = EXCLUDED.source`,
agentID, v.name, v.desc, v.example, v.source); err != nil {
return fmt.Errorf("agent %s var %s: %w", a.key, v.name, err)
}
}
}
// Drop catalog entries for variables that were renamed, so the white-list no
// longer advertises a name templates can't resolve (EngagementTitle→Description).
// 'Now' 从各 agent 目录提升为全局 runtime 变量后,旧库里 goals 仍残留一条 'Now'
// 会与全局项撞名(前端变量列表 key 重复);一并清掉。
if _, err := d.Exec(`DELETE FROM agent_prompt_vars WHERE var_name IN ('EngagementTitle', 'CoverageGaps', 'Now')`); err != nil {
return fmt.Errorf("cleanup renamed vars: %w", err)
}
// Default-on interactive_shell for the runtime agents (planner/worker/mainagent/auto)
// ONCE — respects a later user toggle-off (guarded by a settings flag). goals(one-shot
// decomposer) stays off. Runs after the column exists (schema applied before seed).
if v, _, _ := d.GetSetting("interactive_shell_default_v1"); v != "true" {
if _, err := d.Exec(`UPDATE agents SET interactive_shell=true WHERE key IN ('planner','worker','mainagent','auto')`); err != nil {
return fmt.Errorf("seed interactive_shell defaults: %w", err)
}
_ = d.SetSetting("interactive_shell_default_v1", "true")
}
// Seed the built-in browser (Playwright) MCP once — DISABLED by default (用户
// 需要时自行启用), no proxy by default. The traffic-capture toggle injects/strips
// the recording proxy + CA at runtime (server.Manager.syncBrowserMCPProxy).
// Insert only if absent so we never clobber user edits (args/env/enabled/
// visibility) on restart.
if _, err := d.Exec(`
INSERT INTO mcp_servers(name, transport, command, args, env, enabled)
VALUES ('browser', 'stdio', 'npx', $1, '{}', false)
ON CONFLICT (name) DO NOTHING`,
`["@playwright/mcp","--headless"]`); err != nil {
return fmt.Errorf("seed browser mcp: %w", err)
}
// NOTE: the placeholder ScopeSentry data-source MCP (empty URL + empty X-API-Key,
// disabled) is seeded directly in schema.sql §F so a raw `psql < schema.sql` init
// also gets it. schema.sql is Exec'd on every startup, so it stays idempotent.
if err := d.seedBuiltinSkillVisibility(); err != nil {
return fmt.Errorf("seed skill visibility: %w", err)
}
if err := d.seedDefaultInterceptRules(); err != nil {
return fmt.Errorf("seed intercept rules: %w", err)
}
if err := d.seedDefaultInterceptRulesV2(); err != nil {
return fmt.Errorf("seed intercept rules v2: %w", err)
}
if err := d.seedDefaultInterceptRulesV3(); err != nil {
return fmt.Errorf("seed intercept rules v3: %w", err)
}
if err := d.seedDefaultAssetInterceptRules(); err != nil {
return fmt.Errorf("seed asset intercept rules: %w", err)
}
return nil
}
// seedDefaultAssetInterceptRules inserts the built-in asset blocklist (fuzzy
// domain matches for government / education sites) once on first startup. Gated
// by a settings flag so a user's later disable/delete is never resurrected on
// restart — same policy as the intercept-rule seed.
func (d *DB) seedDefaultAssetInterceptRules() error {
if v, _, _ := d.GetSetting("asset_intercept_default_rules_v1"); v == "done" {
return nil
}
rules := []struct {
kind string
pattern string
note string
}{
{"fuzzy_domain", ".gov", "[内置] 政府网站 (.gov)"},
{"fuzzy_domain", ".gov.cn", "[内置] 政府网站 (.gov.cn)"},
{"fuzzy_domain", ".edu", "[内置] 教育网站 (.edu)"},
{"fuzzy_domain", ".edu.cn", "[内置] 教育网站 (.edu.cn)"},
}
for _, r := range rules {
if _, err := d.Exec(`
INSERT INTO asset_intercept_rules(enabled, kind, pattern, note, builtin)
VALUES (true, $1, $2, $3, true)
ON CONFLICT DO NOTHING`, r.kind, r.pattern, r.note); err != nil {
return fmt.Errorf("asset rule %q: %w", r.pattern, err)
}
}
return d.SetSetting("asset_intercept_default_rules_v1", "done")
}
// builtinSkillVisibility maps a shipped skill's directory name → the built-in
// agent keys that should see it by default. The skill FILES themselves live on the
// filesystem (SkillDir, loaded by norma at runtime); DB only carries this visibility
// binding. Skills omitted here (e.g. playwright-cli, scopesentry) ship invisible by
// default — the user turns them on per-agent when needed. scopesentry additionally
// declares `mcps: ScopeSentry`, which only takes effect once it's made visible and
// that MCP is enabled/configured.
var builtinSkillVisibility = map[string][]string{
"api-recon": {"auto", "pentest", "worker"},
}
// seedBuiltinSkillVisibility binds the shipped built-in skills to their default
// agents. Insert-if-absent (ON CONFLICT DO NOTHING) so a user's later toggle-off is
// never resurrected on restart — matches the browser-MCP / intercept-rule seed policy.
func (d *DB) seedBuiltinSkillVisibility() error {
for skillName, agentKeys := range builtinSkillVisibility {
for _, key := range agentKeys {
if _, err := d.Exec(`
INSERT INTO agent_skill_visibility(agent_id, skill_name, enabled)
SELECT id, $2, true FROM agents WHERE key=$1
ON CONFLICT (agent_id, skill_name) DO NOTHING`, key, skillName); err != nil {
return fmt.Errorf("skill %s → agent %s: %w", skillName, key, err)
}
}
}
return nil
}
// seedDefaultInterceptRules inserts built-in safety intercept rules once on
// first startup. The seed is gated by a settings flag so user edits (disable,
// delete, re-order) are never overwritten on subsequent restarts.
func (d *DB) seedDefaultInterceptRules() error {
if v, _, _ := d.GetSetting("intercept_default_rules_v1"); v == "done" {
return nil
}
type rule struct {
name string
target string // tool_name | tool_input
typ string // string | regex
pattern string
action string
message string
priority int
}
rules := []rule{
// ── 系统破坏性命令 (priority 100) ──────────────────────────────────
{
name: "[内置] 递归强制删除 rm -rf",
target: "tool_input",
typ: "regex",
pattern: `(?i)\brm\b.{0,80}(?:-[a-z]*r[a-z]*f[a-z]*|-[a-z]*f[a-z]*r[a-z]*|--recursive|--no-preserve-root)`,
action: "deny",
message: "禁止执行递归强制删除(rm -rf / rm --recursive),可能永久损坏系统或靶机环境",
priority: 100,
},
{
name: "[内置] 删除系统关键目录",
target: "tool_input",
typ: "regex",
pattern: `\brm\b[^"'\n]{0,60}["'\s](/|/etc|/bin|/usr|/boot|/var|/lib|/sys|/proc|/dev|/sbin|/root)`,
action: "deny",
message: "禁止删除系统关键路径",
priority: 100,
},
{
name: "[内置] 磁盘格式化 mkfs",
target: "tool_input",
typ: "regex",
pattern: `\bmkfs\b`,
action: "deny",
message: "禁止格式化磁盘(mkfs)",
priority: 100,
},
{
name: "[内置] 覆写磁盘设备 dd",
target: "tool_input",
typ: "regex",
pattern: `\bdd\b[^|\n]{0,100}\bof=\s*/dev/[a-zA-Z]`,
action: "deny",
message: "禁止使用 dd 覆写磁盘设备",
priority: 100,
},
{
name: "[内置] Fork 炸弹",
target: "tool_input",
typ: "regex",
pattern: `:\(\)\s*\{[^}]*:\|:`,
action: "deny",
message: "禁止执行 Fork 炸弹",
priority: 100,
},
{
name: "[内置] 关机 / 重启",
target: "tool_input",
typ: "regex",
pattern: `\b(?:shutdown|reboot|halt|poweroff|init\s+[06])\b`,
action: "deny",
message: "禁止执行关机或重启命令",
priority: 100,
},
{
name: "[内置] 杀死全部进程",
target: "tool_input",
typ: "regex",
pattern: `\bkill\s+-9\s+-1\b|\bkillall\s+-9\b`,
action: "deny",
message: "禁止 kill -9 -1 或 killall -9(杀死所有进程)",
priority: 100,
},
{
name: "[内置] 磁盘擦除 shred / wipe",
target: "tool_input",
typ: "regex",
pattern: `\b(?:shred|wipe)\b[^|\n]{0,80}/dev/[a-zA-Z]`,
action: "deny",
message: "禁止对磁盘设备执行 shred/wipe 擦除",
priority: 100,
},
{
name: "[内置] 清空防火墙规则",
target: "tool_input",
typ: "regex",
pattern: `\biptables\s+(?:-F|--flush)\b|\bnft\s+flush\s+ruleset\b`,
action: "deny",
message: "禁止清空防火墙规则(iptables -F / nft flush)",
priority: 100,
},
// ── 数据库破坏性操作 (priority 90) ─────────────────────────────────
{
name: "[内置] SQL DROP DATABASE / TABLE / SCHEMA",
target: "tool_input",
typ: "regex",
pattern: `(?i)\bDROP\s+(?:DATABASE|TABLE|SCHEMA|INDEX|VIEW|TABLESPACE|USER|ROLE)\b`,
action: "deny",
message: "禁止执行 DROP 操作,可能不可逆地销毁数据库对象",
priority: 90,
},
{
name: "[内置] SQL TRUNCATE",
target: "tool_input",
typ: "regex",
pattern: `(?i)\bTRUNCATE\s+(?:TABLE\s+)?\w`,
action: "deny",
message: "禁止执行 TRUNCATE,可能清空数据表所有数据",
priority: 90,
},
{
name: "[内置] MongoDB drop / dropDatabase",
target: "tool_input",
typ: "regex",
pattern: `(?i)\.(?:dropDatabase|dropCollection|drop)\s*\(`,
action: "deny",
message: "禁止执行 MongoDB drop 操作",
priority: 90,
},
{
name: "[内置] Redis FLUSHALL / FLUSHDB",
target: "tool_input",
typ: "regex",
pattern: `(?i)\b(?:FLUSHALL|FLUSHDB)\b`,
action: "deny",
message: "禁止执行 Redis FLUSHALL / FLUSHDB,可能清空全部缓存数据",
priority: 90,
},
// ── HTTP 破坏性请求 (priority 80) ──────────────────────────────────
// Agent 发送 DELETE 请求的三种常见方式:
// 1. curl -X DELETE / --request DELETE(Bash 工具直接执行或写入脚本)
// 2. Python HTTP 客户端 .delete() 方法
// 3. JS/通用脚本里的 method: 'DELETE' / method="DELETE"
{
name: "[内置] curl / wget 发送 DELETE 请求",
target: "tool_input",
typ: "regex",
pattern: `(?i)\bcurl\b[^|\n&;"]{0,300}(?:-X\s*DELETE|--request\s+DELETE|-XDELETE)|\bwget\b[^|\n&;"]{0,300}--method[=\s]+DELETE`,
action: "deny",
message: "禁止通过 curl/wget 发送 HTTP DELETE 请求,可能删除目标系统数据",
priority: 80,
},
{
name: "[内置] Python HTTP 客户端 DELETE(requests/httpx/aiohttp)",
target: "tool_input",
typ: "regex",
pattern: `(?i)\b(?:requests|httpx|aiohttp|urllib\.request)\.delete\s*\(|session\.delete\s*\(|client\.delete\s*\(`,
action: "deny",
message: "禁止使用 Python HTTP 客户端发送 DELETE 请求",
priority: 80,
},
{
name: "[内置] 脚本中声明 HTTP DELETE 方法(JS/通用)",
target: "tool_input",
typ: "regex",
pattern: `(?i)axios\.delete\s*\(|method\s*[:=]\s*['"]DELETE['"]`,
action: "deny",
message: "禁止在脚本中声明并发送 HTTP DELETE 请求",
priority: 80,
},
{
name: "[内置] 批量清空 / 清除接口路径",
target: "tool_input",
typ: "regex",
pattern: `(?i)/(?:clear|wipe|flush|purge|truncate|drop|destroy|factory[-_]reset|reset[-_]all)(?:[/?#"'\s]|$)`,
action: "deny",
message: "禁止调用批量清空或销毁类接口(/clear /wipe /flush /purge 等)",
priority: 80,
},
}
for _, r := range rules {
if _, err := d.Exec(`
INSERT INTO intercept_rules(name, enabled, priority, match_target, match_type, pattern, action, message, timeout_enabled, timeout_seconds, timeout_action)
VALUES ($1, true, $2, $3, $4, $5, $6, $7, false, 60, 'deny')
ON CONFLICT DO NOTHING`,
r.name, r.priority, r.target, r.typ, r.pattern, r.action, r.message,
); err != nil {
return fmt.Errorf("rule %q: %w", r.name, err)
}
}
return d.SetSetting("intercept_default_rules_v1", "done")
}
// seedDefaultInterceptRulesV2 migrates the two safety patterns that used to be
// hard-coded in guard.go (destructive shell + data-exfil pipe) into ordinary
// intercept rules. Gated by its own flag so it also lands on DBs that already ran
// v1. Unlike the old guard.go floor, these are plain [内置] rules — the user can
// disable or delete them. The exfil rule ships DISABLED by default (its
// curl/wget/nc pipe pattern mis-fires on legitimate CTF/pentest reverse-shell and
// data-transfer pipes); enable it manually when exfil gating is actually wanted.
func (d *DB) seedDefaultInterceptRulesV2() error {
if v, _, _ := d.GetSetting("intercept_default_rules_v2"); v == "done" {
return nil
}
rules := []struct {
name string
pattern string
action string
message string
enabled bool
priority int
}{
{
name: "[内置] 破坏性系统命令",
pattern: `(?i)\b(rm\s+-rf\s+/|mkfs|dd\s+if=|:\(\)\s*\{|shutdown|reboot|>\s*/dev/sd)`,
action: "deny",
message: "破坏性命令被拒绝(rm -rf / / mkfs / dd / fork bomb / 关机重启 / 覆写磁盘设备)",
enabled: true,
priority: 100,
},
{
name: "[内置] 数据外泄管道",
pattern: `(?i)(curl|wget|nc|ncat)\b[^|]*\b(\|\s*(curl|wget|nc))`,
action: "deny",
message: "疑似数据外泄管道被拒绝(命令输出经 curl/wget/nc 外传)",
enabled: false,
priority: 80,
},
}
for _, r := range rules {
if _, err := d.Exec(`
INSERT INTO intercept_rules(name, enabled, priority, match_target, match_type, pattern, action, message, timeout_enabled, timeout_seconds, timeout_action)
VALUES ($1, $2, $3, 'tool_input', 'regex', $4, $5, $6, false, 60, 'deny')
ON CONFLICT DO NOTHING`,
r.name, r.enabled, r.priority, r.pattern, r.action, r.message,
); err != nil {
return fmt.Errorf("rule %q: %w", r.name, err)
}
}
return d.SetSetting("intercept_default_rules_v2", "done")
}
// seedDefaultInterceptRulesV3 adds the delete-endpoint path rule. The v1 HTTP rules
// only catch the DELETE *method* (curl -X DELETE, requests.delete(, method:'DELETE'),
// and v1's path rule covers only /clear /wipe /flush /purge /truncate /drop /destroy
// /factory-reset /reset-all — so a plain `curl 'http://t/api/user/delete?id=1'` (a
// delete endpoint reached with GET/POST, which is how most web apps expose deletion)
// slipped through every built-in rule. Own flag so it also lands on DBs that already
// ran v1/v2, where editing the v1 seed would have no effect.
//
// The pattern deliberately requires a separator after the verb so /delivery,
// /details, /delta and /delegate do not match, while /deleteAll, /delete_user and
// /delete-user do. destroy is re-covered here because v1's rule does not allow a
// suffix (/destroyAll was missed).
//
// Exported as a package const only so the seeded regex is unit-testable without a DB.
const deleteEndpointPathPattern = `(?i)/(?:(?:delete|remove|unlink|erase|destroy)[-\w]*|del)(?:[/?#"'\s]|$)`
func (d *DB) seedDefaultInterceptRulesV3() error {
if v, _, _ := d.GetSetting("intercept_default_rules_v3"); v == "done" {
return nil
}
const name = "[内置] 删除类接口路径"
if _, err := d.Exec(`
INSERT INTO intercept_rules(name, enabled, priority, match_target, match_type, pattern, action, message, timeout_enabled, timeout_seconds, timeout_action)
SELECT $1, true, 80, 'tool_input', 'regex', $2, 'deny', $3, false, 60, 'deny'
WHERE NOT EXISTS (SELECT 1 FROM intercept_rules WHERE name = $1)`,
name,
deleteEndpointPathPattern,
"禁止调用删除类接口(/delete /remove /unlink /erase 等),不论使用哪种 HTTP 方法——多数应用的删除接口用 GET/POST 就能触发,同样会真实删除目标数据",
); err != nil {
return fmt.Errorf("rule %q: %w", name, err)
}
return d.SetSetting("intercept_default_rules_v3", "done")
}
+114
View File
@@ -0,0 +1,114 @@
package db
import (
"context"
"database/sql"
"errors"
"testing"
"time"
"github.com/jackc/pgx/v5/pgconn"
)
type schemaExecFunc func(context.Context, string, ...any) (sql.Result, error)
func (f schemaExecFunc) ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) {
return f(ctx, query, args...)
}
func TestApplySchemaRetriesOnlyDeadlocks(t *testing.T) {
deadlockAttempts := 0
delays := []time.Duration{}
err := applySchemaWithRetry(t.Context(), schemaExecFunc(func(context.Context, string, ...any) (sql.Result, error) {
deadlockAttempts++
if deadlockAttempts < 3 {
return nil, &pgconn.PgError{Code: "40P01", Message: "deadlock detected"}
}
return nil, nil
}), func(delay time.Duration) {
delays = append(delays, delay)
})
if err != nil {
t.Fatal(err)
}
if deadlockAttempts != 3 {
t.Fatalf("schema attempts=%d, want 3", deadlockAttempts)
}
if len(delays) != 2 || delays[0] != schemaDeadlockRetryDelays[0] || delays[1] != schemaDeadlockRetryDelays[1] {
t.Fatalf("schema retry delays=%v", delays)
}
nonRetryable := errors.New("permission denied")
nonRetryableAttempts := 0
err = applySchemaWithRetry(t.Context(), schemaExecFunc(func(context.Context, string, ...any) (sql.Result, error) {
nonRetryableAttempts++
return nil, nonRetryable
}), func(time.Duration) {
t.Fatal("non-deadlock error must not be retried")
})
if !errors.Is(err, nonRetryable) || nonRetryableAttempts != 1 {
t.Fatalf("non-retryable result: attempts=%d err=%v", nonRetryableAttempts, err)
}
}
// testDSN returns the configured DSN, skipping the test when neither the env var
// nor a config file supplies one (DSN no longer has a built-in default).
func testDSN(t *testing.T) string {
t.Helper()
dsn, _, err := DSN()
if err != nil {
t.Skipf("no database config (%v) — skipping", err)
}
return dsn
}
// TestOpenSeed opens the live dev PG, applies schema, seeds, and verifies the
// builtin agents + their variable catalog exist. Skips if PG is unreachable.
func TestOpenSeed(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
wantAgents := len(builtinAgents)
var agents int
if err := d.QueryRow(`SELECT count(*) FROM agents WHERE builtin`).Scan(&agents); err != nil {
t.Fatal(err)
}
if agents != wantAgents {
t.Fatalf("want %d builtin agents, got %d", wantAgents, agents)
}
// planner should have at least its seeded catalog vars
wantPlannerVars := 0
for _, a := range builtinAgents {
if a.key == "planner" {
wantPlannerVars = len(a.vars)
break
}
}
var plannerVars int
if err := d.QueryRow(`
SELECT count(*) FROM agent_prompt_vars v
JOIN agents a ON a.id = v.agent_id
WHERE a.key = 'planner'`).Scan(&plannerVars); err != nil {
t.Fatal(err)
}
if plannerVars < wantPlannerVars {
t.Fatalf("want at least %d planner vars, got %d", wantPlannerVars, plannerVars)
}
// re-open must be idempotent (no duplicate agents)
d2, err := Open(testDSN(t))
if err != nil {
t.Fatal(err)
}
defer d2.Close()
if err := d2.QueryRow(`SELECT count(*) FROM agents WHERE builtin`).Scan(&agents); err != nil {
t.Fatal(err)
}
if agents != wantAgents {
t.Fatalf("after re-open want %d agents, got %d", wantAgents, agents)
}
}
+253
View File
@@ -0,0 +1,253 @@
package db
// cold-digest §1/§2.3/§5: persistence for cold-node compression — the round
// counter, cold_since_round stamps, content versions, digest nodes and the
// covers edges that are the source of truth for "which digest folds node X".
import (
"database/sql"
"encoding/json"
"fmt"
"sort"
"strings"
)
// BumpRound advances this exploration's planner-round counter and returns the
// new value (§2.3). Called once per planner wake-up.
func (s *ExplorationStore) BumpRound() (int64, error) {
var r int64
err := s.db.QueryRow(`UPDATE explorations SET round_no = round_no + 1 WHERE id=$1 RETURNING round_no`, s.expID).Scan(&r)
return r, err
}
// RoundNo returns the current planner-round counter.
func (s *ExplorationStore) RoundNo() (int64, error) {
var r int64
err := s.db.QueryRow(`SELECT round_no FROM explorations WHERE id=$1`, s.expID).Scan(&r)
return r, err
}
// ColdStamps returns cold_since_round for every foldable node (intent/fact):
// id → *round (nil when the node is hot / unstamped). Used to compute the ≥R
// debounce and to know which stamps to set/clear this round.
func (s *ExplorationStore) ColdStamps() (map[int64]*int64, error) {
rows, err := s.db.Query(`SELECT id, cold_since_round FROM exploration_nodes
WHERE exploration_id=$1 AND kind IN ('intent','fact')`, s.expID)
if err != nil {
return nil, err
}
defer rows.Close()
out := map[int64]*int64{}
for rows.Next() {
var id int64
var cs sql.NullInt64
if err := rows.Scan(&id, &cs); err != nil {
return nil, err
}
if cs.Valid {
v := cs.Int64
out[id] = &v
} else {
out[id] = nil
}
}
return out, rows.Err()
}
// StampOp is one cold_since_round change: Set stamps Round; !Set clears it.
type StampOp struct {
ID int64
Set bool
Round int64
}
// ApplyStampOps writes a batch of cold_since_round changes in one transaction.
func (s *ExplorationStore) ApplyStampOps(ops []StampOp) error {
if len(ops) == 0 {
return nil
}
tx, err := s.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
for _, o := range ops {
if o.Set {
if _, err := tx.Exec(`UPDATE exploration_nodes SET cold_since_round=$1 WHERE id=$2 AND exploration_id=$3`, o.Round, o.ID, s.expID); err != nil {
return err
}
} else {
if _, err := tx.Exec(`UPDATE exploration_nodes SET cold_since_round=NULL WHERE id=$1 AND exploration_id=$2`, o.ID, s.expID); err != nil {
return err
}
}
}
return tx.Commit()
}
// ContentVersions returns id → content_version for all nodes (§5.3 signature).
func (s *ExplorationStore) ContentVersions() (map[int64]int, error) {
rows, err := s.db.Query(`SELECT id, content_version FROM exploration_nodes WHERE exploration_id=$1`, s.expID)
if err != nil {
return nil, err
}
defer rows.Close()
out := map[int64]int{}
for rows.Next() {
var id int64
var v int
if err := rows.Scan(&id, &v); err != nil {
return nil, err
}
out[id] = v
}
return out, rows.Err()
}
// ActiveDigests returns the exploration's live digest nodes (state='active'),
// oldest first.
func (s *ExplorationStore) ActiveDigests() ([]*Node, error) {
rows, err := s.db.Query(`SELECT `+nodeCols+` FROM exploration_nodes
WHERE exploration_id=$1 AND kind=$2 AND state=$3 ORDER BY id`, s.expID, KindDigest, StateDigestActive)
if err != nil {
return nil, err
}
defer rows.Close()
return scanNodes(rows)
}
// AddDigest writes one digest node and its covers edges (digest→member) in a
// single transaction. payload is the digest body + member_ids + generation +
// signature (see cold-digest §1). Returns the new digest id.
func (s *ExplorationStore) AddDigest(payload map[string]any, memberIDs []int64) (int64, error) {
raw, _ := json.Marshal(payload)
tx, err := s.db.Begin()
if err != nil {
return 0, err
}
defer tx.Rollback()
var id int64
if err := tx.QueryRow(`
INSERT INTO exploration_nodes(exploration_id, kind, payload, priority, state, origin)
VALUES ($1, $2, $3, 0, $4, 'compactor') RETURNING id`,
s.expID, KindDigest, string(raw), StateDigestActive).Scan(&id); err != nil {
return 0, err
}
for _, m := range memberIDs {
if m == id {
continue
}
if _, err := tx.Exec(`
INSERT INTO exploration_edges(exploration_id, src_id, rel, dst_id) VALUES ($1,$2,$3,$4)
ON CONFLICT (exploration_id, src_id, rel, dst_id) DO NOTHING`, s.expID, id, RelCovers, m); err != nil {
return 0, err
}
}
return id, tx.Commit()
}
// CoveredMembers maps member id → covering digest id, for ACTIVE digests only
// (§6.1 point 1). A node with no entry is not currently folded. Should a member
// carry covers edges from two digests (a torn major write), the lowest digest id
// wins deterministically — callers dedupe on this.
func (s *ExplorationStore) CoveredMembers() (map[int64]int64, error) {
rows, err := s.db.Query(`
SELECT e.dst_id, e.src_id
FROM exploration_edges e
JOIN exploration_nodes d ON d.id=e.src_id AND d.exploration_id=e.exploration_id
WHERE e.exploration_id=$1 AND e.rel=$2 AND d.kind=$3 AND d.state=$4
ORDER BY e.src_id`, s.expID, RelCovers, KindDigest, StateDigestActive)
if err != nil {
return nil, err
}
defer rows.Close()
out := map[int64]int64{}
for rows.Next() {
var member, digest int64
if err := rows.Scan(&member, &digest); err != nil {
return nil, err
}
if _, seen := out[member]; !seen { // first (lowest digest id) wins
out[member] = digest
}
}
return out, rows.Err()
}
// DigestMembers returns the member ids a digest covers (covers edges), sorted.
func (s *ExplorationStore) DigestMembers(digestID int64) ([]int64, error) {
rows, err := s.db.Query(`SELECT dst_id FROM exploration_edges
WHERE exploration_id=$1 AND src_id=$2 AND rel=$3`, s.expID, digestID, RelCovers)
if err != nil {
return nil, err
}
defer rows.Close()
var ids []int64
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
return nil, err
}
ids = append(ids, id)
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
return ids, rows.Err()
}
// SupersedeDigests retires digest nodes (state→superseded) AND removes their
// covers edges, atomically, so the active-coverage set (CoveredMembers) never
// double-counts a member during a major recompaction (§5.1). The digest node
// itself is kept (node_detail can still resolve it).
func (s *ExplorationStore) SupersedeDigests(ids []int64) error {
if len(ids) == 0 {
return nil
}
tx, err := s.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
for _, id := range ids {
if _, err := tx.Exec(`DELETE FROM exploration_edges
WHERE exploration_id=$1 AND src_id=$2 AND rel=$3`, s.expID, id, RelCovers); err != nil {
return err
}
if _, err := tx.Exec(`UPDATE exploration_nodes SET state=$1 WHERE id=$2 AND exploration_id=$3 AND kind=$4`,
StateDigestSuperseded, id, s.expID, KindDigest); err != nil {
return err
}
}
return tx.Commit()
}
// NodeAssets maps each of the given node ids → the asset ids it is anchored to
// (exploration_anchors). Used to group cold digests by asset (§6.2 index).
func (s *ExplorationStore) NodeAssets(ids []int64) (map[int64][]int64, error) {
out := map[int64][]int64{}
if len(ids) == 0 {
return out, nil
}
placeholders := make([]string, len(ids))
args := make([]any, 0, len(ids)+1)
args = append(args, s.expID)
for i, id := range ids {
placeholders[i] = fmt.Sprintf("$%d", i+2)
args = append(args, id)
}
rows, err := s.db.Query(`SELECT a.node_id, a.asset_id
FROM exploration_anchors a
JOIN exploration_nodes n ON n.id=a.node_id
WHERE n.exploration_id=$1 AND a.node_id IN (`+strings.Join(placeholders, ",")+`)`, args...)
if err != nil {
return nil, err
}
defer rows.Close()
for rows.Next() {
var node, asset int64
if err := rows.Scan(&node, &asset); err != nil {
return nil, err
}
out[node] = append(out[node], asset)
}
return out, rows.Err()
}
+2021
View File
File diff suppressed because it is too large Load Diff
+397
View File
@@ -0,0 +1,397 @@
package db
import (
"database/sql"
"sort"
)
// DirectSourceStore binds one directly related task to its exploration store.
// It is intentionally a read-side helper: callers keep using the receiver store
// for every graph mutation, frontier lookup, and intent claim.
type DirectSourceStore struct {
Task TaskSource
Store *ExplorationStore
}
// DirectSourceStores resolves the live, direct task relations for this
// exploration. It deliberately queries on every call: relations disappear when
// a source task is deleted, and inherited context must not retain stale rows.
// Sources of a source are never expanded.
func (s *ExplorationStore) DirectSourceStores() ([]DirectSourceStore, error) {
rows, err := s.db.Query(`
SELECT source.id, source.exploration_id, source.description, source.goal, source.status
FROM tasks owner
JOIN task_relations relation ON relation.task_id=owner.id
JOIN tasks source ON source.id=relation.source_task_id AND source.deleted_at IS NULL
WHERE owner.exploration_id=$1 AND owner.deleted_at IS NULL
ORDER BY relation.created_at, source.id`, s.expID)
if err != nil {
return nil, err
}
defer rows.Close()
out := []DirectSourceStore{}
for rows.Next() {
var source TaskSource
if err := rows.Scan(&source.TaskID, &source.ExplorationID, &source.Description, &source.Goal, &source.Status); err != nil {
return nil, err
}
out = append(out, DirectSourceStore{Task: source, Store: s.db.Exploration(source.ExplorationID)})
}
return out, rows.Err()
}
// TaskID returns the live task bound to this exploration. Explorations created
// directly in tests or maintenance code have no task and return zero.
func (s *ExplorationStore) TaskID() (int64, error) {
var id int64
err := s.db.QueryRow(`SELECT id FROM tasks WHERE exploration_id=$1 AND deleted_at IS NULL`, s.expID).Scan(&id)
if err != nil {
if err == sql.ErrNoRows {
return 0, nil
}
return 0, err
}
return id, nil
}
func markInheritedNode(n *Node, taskID int64) *Node {
if n == nil {
return nil
}
n.SourceTaskID = taskID
n.Inherited = true
return n
}
func markInheritedActivities(in []Activity, taskID int64) []Activity {
for i := range in {
in[i].SourceTaskID = taskID
in[i].Inherited = true
}
return in
}
func inheritedIntentTerminal(state string) bool {
switch state {
case "done", "blocked", "exhausted", "stopped":
return true
default:
return false
}
}
// ListByKindWithSources returns local nodes followed by nodes from each direct
// source in relation order. The limit remains per exploration, matching the
// existing ListByKind contract while ensuring one large task cannot hide all
// inherited context from another source.
func (s *ExplorationStore) ListByKindWithSources(kind string, limit int) ([]*Node, error) {
out, err := s.ListByKind(kind, limit)
if err != nil {
return nil, err
}
sources, err := s.DirectSourceStores()
if err != nil {
return nil, err
}
for _, source := range sources {
nodes, err := source.Store.ListByKind(kind, limit)
if err != nil {
return nil, err
}
for _, node := range nodes {
if kind == KindIntent && !inheritedIntentTerminal(node.State) {
continue
}
out = append(out, markInheritedNode(node, source.Task.TaskID))
}
}
return out, nil
}
// ListByKindPageWithSources is the paginated, keyword-filterable sibling of
// ListByKindWithSources: it returns one newest-first page (id < before, before<=0
// = newest) spanning this exploration and its direct sources, plus hasMore and the
// filtered total across all of them. Node ids are globally unique, so merging each
// store's own page and re-sorting by id DESC yields the true global page; fetching
// limit+1 per store guarantees the merged top-`limit` is complete.
func (s *ExplorationStore) ListByKindPageWithSources(kind string, before int64, limit int, q string) (nodes []*Node, hasMore bool, total int, err error) {
if limit <= 0 {
limit = 20
}
own, err := s.listByKindPageFiltered(kind, before, limit, q)
if err != nil {
return nil, false, 0, err
}
merged := own
total, err = s.countByKindFiltered(kind, q)
if err != nil {
return nil, false, 0, err
}
sources, err := s.DirectSourceStores()
if err != nil {
return nil, false, 0, err
}
for _, source := range sources {
page, err := source.Store.listByKindPageFiltered(kind, before, limit, q)
if err != nil {
return nil, false, 0, err
}
for _, n := range page {
merged = append(merged, markInheritedNode(n, source.Task.TaskID))
}
cnt, err := source.Store.countByKindFiltered(kind, q)
if err != nil {
return nil, false, 0, err
}
total += cnt
}
sort.Slice(merged, func(i, j int) bool { return merged[i].ID > merged[j].ID })
hasMore = len(merged) > limit
if hasMore {
merged = merged[:limit]
}
return merged, hasMore, total, nil
}
// GetNodeWithSources reads a node only when it belongs to this exploration or
// one of its direct sources. Inherited nodes are tagged so tool callers can keep
// them read-only and show their provenance.
func (s *ExplorationStore) GetNodeWithSources(id int64) (*Node, error) {
node, err := s.GetNode(id)
if err != nil || node != nil {
return node, err
}
sources, err := s.DirectSourceStores()
if err != nil {
return nil, err
}
for _, source := range sources {
node, err = source.Store.GetNode(id)
if err != nil {
return nil, err
}
if node != nil {
if node.Kind == KindIntent && !inheritedIntentTerminal(node.State) {
continue
}
return markInheritedNode(node, source.Task.TaskID), nil
}
}
return nil, nil
}
// FindingIntentsWithSources combines finding lineage for the current
// exploration and each direct source. Node ids are globally unique.
func (s *ExplorationStore) FindingIntentsWithSources() (map[int64]int64, error) {
out, err := s.FindingIntents()
if err != nil {
return nil, err
}
sources, err := s.DirectSourceStores()
if err != nil {
return nil, err
}
for _, source := range sources {
items, err := source.Store.FindingIntentsTerminal()
if err != nil {
return nil, err
}
for findingID, intentID := range items {
out[findingID] = intentID
}
}
return out, nil
}
// ActivityTraceWithSources returns a work trace when its intent belongs to the
// current exploration or a direct source. It never searches indirect sources.
func (s *ExplorationStore) ActivityTraceWithSources(nodeID int64, limit int) ([]Activity, error) {
acts, err := s.ActivityTrace(nodeID, limit)
if err != nil || len(acts) > 0 {
return acts, err
}
sources, err := s.DirectSourceStores()
if err != nil {
return nil, err
}
for _, source := range sources {
node, nodeErr := source.Store.GetNode(nodeID)
if nodeErr != nil {
return nil, nodeErr
}
if node == nil || node.Kind != KindIntent {
continue
}
acts, err = source.Store.ActivityTraceForTerminalIntent(nodeID, limit)
if err != nil {
return nil, err
}
if len(acts) > 0 {
return markInheritedActivities(acts, source.Task.TaskID), nil
}
}
return []Activity{}, nil
}
// ActivityListWithSources is the source-aware equivalent used by
// get_worker_output. Node ids are global, so the first owning exploration is
// unambiguous even when the work has not emitted any activity yet.
func (s *ExplorationStore) ActivityListWithSources(nodeID, sinceID int64, limit int) ([]Activity, int64, error) {
node, err := s.GetNode(nodeID)
if err != nil {
return nil, sinceID, err
}
if node != nil {
return s.ActivityList(&nodeID, sinceID, limit)
}
sources, err := s.DirectSourceStores()
if err != nil {
return nil, sinceID, err
}
for _, source := range sources {
node, err = source.Store.GetNode(nodeID)
if err != nil {
return nil, sinceID, err
}
if node == nil || node.Kind != KindIntent {
continue
}
acts, cursor, err := source.Store.ActivityListForTerminalIntent(nodeID, sinceID, limit)
if err != nil {
return nil, sinceID, err
}
return markInheritedActivities(acts, source.Task.TaskID), cursor, nil
}
return []Activity{}, sinceID, nil
}
// ActivityDetailWithSources keeps the legacy local-task lookup (including local
// thinking rows), while inherited details are restricted to terminal worker
// intents. Source planner/main rows have no node id and must never become part of
// inherited context.
func (s *ExplorationStore) ActivityDetailWithSources(id int64) (string, error) {
detail, err := s.ActivityDetail(id)
if err != nil || detail != "" {
return detail, err
}
acts, err := s.ActivityByIDsWithSources([]int64{id})
if err != nil || len(acts) == 0 {
return "", err
}
return acts[0].Detail, nil
}
// ActivityTraceSearchWithSources performs the scoped worker-trace search used
// by get_worker_trace. A node id resolves to at most one exploration because
// exploration node ids are global.
func (s *ExplorationStore) ActivityTraceSearchWithSources(nodeID int64, q string, limit int) ([]Activity, error) {
acts, err := s.ActivityTraceSearch(&nodeID, q, limit)
if err != nil || len(acts) > 0 {
return acts, err
}
sources, err := s.DirectSourceStores()
if err != nil {
return nil, err
}
for _, source := range sources {
node, nodeErr := source.Store.GetNode(nodeID)
if nodeErr != nil {
return nil, nodeErr
}
if node == nil || node.Kind != KindIntent {
continue
}
acts, err = source.Store.ActivityTraceSearchForTerminalIntent(nodeID, q, limit)
if err != nil {
return nil, err
}
if len(acts) > 0 {
return markInheritedActivities(acts, source.Task.TaskID), nil
}
}
return []Activity{}, nil
}
// ActivityTraceSearchAllWithSources searches local worker traces plus every
// direct source. The local owner is excluded only from the current exploration;
// inherited traces are immutable historical context.
func (s *ExplorationStore) ActivityTraceSearchAllWithSources(excludeNodeID int64, q string, limit int) ([]Activity, error) {
if limit <= 0 {
limit = 100
}
out, err := s.ActivityTraceSearchExcluding(excludeNodeID, q, limit)
if err != nil {
return nil, err
}
sources, err := s.DirectSourceStores()
if err != nil {
return nil, err
}
for _, source := range sources {
acts, err := source.Store.ActivityTraceSearchTerminalIntents(q, limit)
if err != nil {
return nil, err
}
out = append(out, markInheritedActivities(acts, source.Task.TaskID)...)
}
sort.Slice(out, func(i, j int) bool { return out[i].ID < out[j].ID })
if len(out) > limit {
out = out[:limit]
}
return out, nil
}
// ActivityByIDsWithSources loads local step details with the legacy behavior. For
// direct sources it only returns rows attached to terminal intents, preventing an
// arbitrary global activity id from exposing source planner/main transcripts.
func (s *ExplorationStore) ActivityByIDsWithSources(ids []int64) ([]Activity, error) {
out, err := s.ActivityByIDs(ids)
if err != nil {
return nil, err
}
sources, err := s.DirectSourceStores()
if err != nil {
return nil, err
}
for _, source := range sources {
acts, err := source.Store.ActivityByIDsForTerminalIntents(ids)
if err != nil {
return nil, err
}
out = append(out, markInheritedActivities(acts, source.Task.TaskID)...)
}
sort.Slice(out, func(i, j int) bool { return out[i].ID < out[j].ID })
return out, nil
}
// AssetRefsWithSources returns anchored nodes from this exploration and each
// direct source. Inherited entries retain their owning task id so API/UI callers
// can present them as immutable context.
func (s *ExplorationStore) AssetRefsWithSources(assetID int64) ([]AssetRef, error) {
out, err := s.AssetRefs(assetID)
if err != nil {
return nil, err
}
sources, err := s.DirectSourceStores()
if err != nil {
return nil, err
}
for _, source := range sources {
refs, err := source.Store.AssetRefs(assetID)
if err != nil {
return nil, err
}
for i := range refs {
if refs[i].Kind == KindIntent && !inheritedIntentTerminal(refs[i].State) {
continue
}
refs[i].SourceTaskID = source.Task.TaskID
refs[i].Inherited = true
out = append(out, refs[i])
}
}
sort.Slice(out, func(i, j int) bool { return out[i].ID > out[j].ID })
return out, nil
}
+431
View File
@@ -0,0 +1,431 @@
package db
import (
"fmt"
"testing"
)
func TestInheritedActivityReadsRequireTerminalIntent(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
expID, err := d.CreateExploration("source activity boundary", "source activity boundary")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _, _ = d.Exec(`DELETE FROM explorations WHERE id=$1`, expID) })
store := d.Exploration(expID)
intentID, err := store.AddIntent(map[string]any{"summary": "source work"}, 1, nil, "planner")
if err != nil {
t.Fatal(err)
}
appendStep := func(kind, summary string) int64 {
t.Helper()
id, appendErr := store.AppendActivity(Activity{
NodeID: &intentID, Worker: "worker", Kind: kind, Summary: summary, Detail: summary + " detail",
})
if appendErr != nil {
t.Fatal(appendErr)
}
return id
}
textID := appendStep("text", "shared text")
thinkingID := appendStep("thinking", "private reasoning")
usageID := appendStep("usage", "private accounting")
resultID := appendStep("result", "shared result")
if err := store.SetIntentState(intentID, "done"); err != nil {
t.Fatal(err)
}
page, more, err := store.ActivityPageForTerminalIntent(intentID, 0, 10)
if err != nil || more || len(page) != 2 || page[0].ID != textID || page[1].ID != resultID {
t.Fatalf("terminal page boundary: page=%+v more=%v err=%v", page, more, err)
}
list, cursor, err := store.ActivityListForTerminalIntent(intentID, 0, 10)
if err != nil || len(list) != 2 || cursor != resultID {
t.Fatalf("terminal list boundary: list=%+v cursor=%d err=%v", list, cursor, err)
}
trace, err := store.ActivityTraceForTerminalIntent(intentID, 10)
if err != nil || len(trace) != 2 {
t.Fatalf("terminal trace boundary: trace=%+v err=%v", trace, err)
}
hits, err := store.ActivityTraceSearchForTerminalIntent(intentID, "shared", 10)
if err != nil || len(hits) != 2 {
t.Fatalf("terminal search boundary: hits=%+v err=%v", hits, err)
}
details, err := store.ActivityByIDsForTerminalIntents([]int64{textID, thinkingID, usageID, resultID})
if err != nil || len(details) != 2 || details[0].ID != textID || details[1].ID != resultID {
t.Fatalf("terminal detail boundary: details=%+v err=%v", details, err)
}
// Reopening a source intent must close every inherited activity read even if
// callers still hold a stale terminal-state snapshot.
if err := store.SetIntentState(intentID, "running"); err != nil {
t.Fatal(err)
}
page, _, err = store.ActivityPageForTerminalIntent(intentID, 0, 10)
if err != nil || len(page) != 0 {
t.Fatalf("reopened intent leaked through page: page=%+v err=%v", page, err)
}
list, _, err = store.ActivityListForTerminalIntent(intentID, 0, 10)
if err != nil || len(list) != 0 {
t.Fatalf("reopened intent leaked through list: list=%+v err=%v", list, err)
}
trace, err = store.ActivityTraceForTerminalIntent(intentID, 10)
if err != nil || len(trace) != 0 {
t.Fatalf("reopened intent leaked through trace: trace=%+v err=%v", trace, err)
}
hits, err = store.ActivityTraceSearchForTerminalIntent(intentID, "shared", 10)
if err != nil || len(hits) != 0 {
t.Fatalf("reopened intent leaked through search: hits=%+v err=%v", hits, err)
}
details, err = store.ActivityByIDsForTerminalIntents([]int64{textID, resultID})
if err != nil || len(details) != 0 {
t.Fatalf("reopened intent leaked through detail: details=%+v err=%v", details, err)
}
}
func TestExplorationDirectSourceReadView(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
grand, err := d.CreateTask("grand source", "grand goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
source, err := d.CreateTaskWithOptions("direct source", "source goal", TaskCreateOptions{
SourceTaskIDs: []int64{grand.ID},
})
if err != nil {
t.Fatal(err)
}
current, err := d.CreateTaskWithOptions("current", "current goal", TaskCreateOptions{
SourceTaskIDs: []int64{source.ID},
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
_ = d.DeleteTask(current.ID)
_ = d.DeleteTask(source.ID)
_ = d.DeleteTask(grand.ID)
})
grandStore := d.Exploration(grand.ExplorationID)
sourceStore := d.Exploration(source.ExplorationID)
currentStore := d.Exploration(current.ExplorationID)
grandFact, err := grandStore.AddNode(KindFact, map[string]any{"summary": "grand-only fact"}, 0, "confirmed", "worker", nil)
if err != nil {
t.Fatal(err)
}
sourceFact, err := sourceStore.AddNode(KindFact, map[string]any{"summary": "direct source fact"}, 0, "confirmed", "worker", nil)
if err != nil {
t.Fatal(err)
}
sourceFinding, err := sourceStore.AddNode(KindFinding, map[string]any{"summary": "direct source finding"}, 0, "confirmed", "worker", nil)
if err != nil {
t.Fatal(err)
}
sourceIntent, err := sourceStore.AddIntent(map[string]any{"summary": "source work"}, 9, nil, "planner")
if err != nil {
t.Fatal(err)
}
if err := sourceStore.Link(sourceIntent, RelYields, sourceFinding); err != nil {
t.Fatal(err)
}
currentIntent, err := currentStore.AddIntent(map[string]any{"summary": "current work"}, 1, nil, "planner")
if err != nil {
t.Fatal(err)
}
currentFact, err := currentStore.AddNode(KindFact, map[string]any{"summary": "current fact"}, 0, "confirmed", "worker", nil)
if err != nil {
t.Fatal(err)
}
sources, err := currentStore.DirectSourceStores()
if err != nil {
t.Fatal(err)
}
if len(sources) != 1 || sources[0].Task.TaskID != source.ID {
t.Fatalf("want only direct source %d, got %+v", source.ID, sources)
}
frontier, err := currentStore.Frontier(20)
if err != nil {
t.Fatal(err)
}
if len(frontier) != 1 || frontier[0].ID != currentIntent {
t.Fatalf("source intent leaked into local frontier: %+v", frontier)
}
if claimed, err := currentStore.ClaimIntent(sourceIntent, "wrong-task-worker"); err != nil || claimed {
t.Fatalf("source intent must not be claimable: claimed=%v err=%v", claimed, err)
}
if got, err := currentStore.GetNodeWithSources(sourceIntent); err != nil || got != nil {
t.Fatalf("live source intent must not be readable as inherited history: node=%+v err=%v", got, err)
}
if got, err := currentStore.ListByKindWithSources(KindIntent, 100); err != nil || len(got) != 1 || got[0].ID != currentIntent {
t.Fatalf("live source intent leaked through inherited list: nodes=%+v err=%v", got, err)
}
if got, err := currentStore.FindingIntentsWithSources(); err != nil || got[sourceFinding] != 0 {
t.Fatalf("live source intent leaked through finding lineage: lineage=%+v err=%v", got, err)
}
facts, err := currentStore.ListByKindWithSources(KindFact, 100)
if err != nil {
t.Fatal(err)
}
seenCurrent, seenSource, seenGrand := false, false, false
for _, fact := range facts {
switch fact.ID {
case currentFact:
seenCurrent = !fact.Inherited && fact.SourceTaskID == 0
case sourceFact:
seenSource = fact.Inherited && fact.SourceTaskID == source.ID
case grandFact:
seenGrand = true
}
}
if !seenCurrent || !seenSource || seenGrand {
t.Fatalf("unexpected direct-source fact view: current=%v source=%v grand=%v", seenCurrent, seenSource, seenGrand)
}
gotSource, err := currentStore.GetNodeWithSources(sourceFact)
if err != nil || gotSource == nil || !gotSource.Inherited || gotSource.SourceTaskID != source.ID {
t.Fatalf("source node lookup: node=%+v err=%v", gotSource, err)
}
gotGrand, err := currentStore.GetNodeWithSources(grandFact)
if err != nil || gotGrand != nil {
t.Fatalf("indirect source must be invisible: node=%+v err=%v", gotGrand, err)
}
// Existing write methods remain bound to currentStore.expID. Passing an
// inherited id is a no-op and cannot mutate the source blackboard.
if err := currentStore.SetNodeState(sourceFact, "dismissed"); err != nil {
t.Fatal(err)
}
unchanged, err := sourceStore.GetNode(sourceFact)
if err != nil || unchanged == nil || unchanged.State != "confirmed" {
t.Fatalf("inherited fact was mutated: node=%+v err=%v", unchanged, err)
}
if err := sourceStore.SetIntentState(sourceIntent, "done"); err != nil {
t.Fatal(err)
}
if got, err := currentStore.GetNodeWithSources(sourceIntent); err != nil || got == nil || !got.Inherited {
t.Fatalf("terminal source intent should become readable history: node=%+v err=%v", got, err)
}
if got, err := currentStore.FindingIntentsWithSources(); err != nil || got[sourceFinding] != sourceIntent {
t.Fatalf("terminal source finding lineage missing: lineage=%+v err=%v", got, err)
}
sourcePlannerStep, err := sourceStore.AppendActivity(Activity{
Worker: "planner", Kind: "text", Summary: "private source planner", Detail: "private source planner transcript",
})
if err != nil {
t.Fatal(err)
}
openSourceIntent, err := sourceStore.AddIntent(map[string]any{"summary": "source work still running"}, 1, nil, "planner")
if err != nil {
t.Fatal(err)
}
if err := sourceStore.SetIntentState(openSourceIntent, "running"); err != nil {
t.Fatal(err)
}
openSourceStep, err := sourceStore.AppendActivity(Activity{
NodeID: &openSourceIntent, Worker: "source-live-worker", Kind: "text",
Summary: "private live source work", Detail: "private live source transcript",
})
if err != nil {
t.Fatal(err)
}
sourceStep, err := sourceStore.AppendActivity(Activity{
NodeID: &sourceIntent, Worker: "source-worker", Kind: "tool_result", Tool: "HTTP",
Summary: "source-secret response", Detail: "source-secret full detail",
})
if err != nil {
t.Fatal(err)
}
currentPlannerStep, err := currentStore.AppendActivity(Activity{
Worker: "planner", Kind: "text", Summary: "current planner", Detail: "current planner detail",
})
if err != nil {
t.Fatal(err)
}
currentThinkingStep, err := currentStore.AppendActivity(Activity{
Worker: "planner", Kind: "thinking", Summary: "current thinking", Detail: "current thinking detail",
})
if err != nil {
t.Fatal(err)
}
grandStep, err := grandStore.AppendActivity(Activity{
NodeID: &grandFact, Worker: "grand-worker", Kind: "tool_result", Tool: "HTTP",
Summary: "grand-only response", Detail: "grand-only detail",
})
if err != nil {
t.Fatal(err)
}
trace, err := currentStore.ActivityTraceWithSources(sourceIntent, 20)
if err != nil || len(trace) != 1 || trace[0].ID != sourceStep || !trace[0].Inherited || trace[0].SourceTaskID != source.ID {
t.Fatalf("source trace lookup: trace=%+v err=%v", trace, err)
}
hits, err := currentStore.ActivityTraceSearchAllWithSources(0, "source-secret", 20)
if err != nil || len(hits) != 1 || hits[0].ID != sourceStep || !hits[0].Inherited {
t.Fatalf("source trace search: hits=%+v err=%v", hits, err)
}
grandHits, err := currentStore.ActivityTraceSearchAllWithSources(0, "grand-only", 20)
if err != nil || len(grandHits) != 0 {
t.Fatalf("indirect source trace leaked: hits=%+v err=%v", grandHits, err)
}
liveHits, err := currentStore.ActivityTraceSearchAllWithSources(0, "private live source", 20)
if err != nil || len(liveHits) != 0 {
t.Fatalf("live source trace leaked: hits=%+v err=%v", liveHits, err)
}
if liveTrace, err := currentStore.ActivityTraceWithSources(openSourceIntent, 20); err != nil || len(liveTrace) != 0 {
t.Fatalf("live inherited intent trace leaked: trace=%+v err=%v", liveTrace, err)
}
details, err := currentStore.ActivityByIDsWithSources([]int64{
sourcePlannerStep, openSourceStep, sourceStep, currentPlannerStep, grandStep,
})
if err != nil || len(details) != 2 || details[0].ID != sourceStep || !details[0].Inherited || details[1].ID != currentPlannerStep || details[1].Inherited {
t.Fatalf("source trace detail allow-list: details=%+v err=%v", details, err)
}
if detail, err := currentStore.ActivityDetailWithSources(sourcePlannerStep); err != nil || detail != "" {
t.Fatalf("source planner transcript leaked: detail=%q err=%v", detail, err)
}
if detail, err := currentStore.ActivityDetailWithSources(openSourceStep); err != nil || detail != "" {
t.Fatalf("live source intent transcript leaked: detail=%q err=%v", detail, err)
}
if detail, err := currentStore.ActivityDetailWithSources(currentPlannerStep); err != nil || detail != "current planner detail" {
t.Fatalf("local activity detail behavior changed: detail=%q err=%v", detail, err)
}
if detail, err := currentStore.ActivityDetailWithSources(currentThinkingStep); err != nil || detail != "current thinking detail" {
t.Fatalf("local thinking detail behavior changed: detail=%q err=%v", detail, err)
}
}
func TestTaskAssetContextUsesDirectSourceScopeAndAnchors(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
grand, err := d.CreateTask("grand assets", "grand goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
source, err := d.CreateTaskWithOptions("source assets", "source goal", TaskCreateOptions{SourceTaskIDs: []int64{grand.ID}})
if err != nil {
t.Fatal(err)
}
current, err := d.CreateTaskWithOptions("current assets", "current goal", TaskCreateOptions{SourceTaskIDs: []int64{source.ID}})
if err != nil {
t.Fatal(err)
}
assets := d.Assets()
sourceTestedDomain := fmt.Sprintf("context-tested-%d.invalid", source.ID)
sourceUntestedDomain := fmt.Sprintf("context-untested-%d.invalid", source.ID)
sourceAnchoredOnlyDomain := fmt.Sprintf("context-anchor-only-%d.invalid", source.ID)
grandDomain := fmt.Sprintf("context-grand-%d.invalid", grand.ID)
sourceTestedID, err := assets.UpsertRootDomain(UpsertRootDomainReq{Domain: sourceTestedDomain, TaskID: source.ID})
if err != nil {
t.Fatal(err)
}
sourceUntestedID, err := assets.UpsertRootDomain(UpsertRootDomainReq{Domain: sourceUntestedDomain, TaskID: source.ID})
if err != nil {
t.Fatal(err)
}
sourceAnchoredOnlyID, err := assets.UpsertRootDomain(UpsertRootDomainReq{Domain: sourceAnchoredOnlyDomain, TaskID: source.ID})
if err != nil {
t.Fatal(err)
}
grandID, err := assets.UpsertRootDomain(UpsertRootDomainReq{Domain: grandDomain, TaskID: grand.ID})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
_ = d.DeleteTask(current.ID)
_ = d.DeleteTask(source.ID)
_ = d.DeleteTask(grand.ID)
_, _ = assets.DeleteByIDs([]int64{sourceTestedID, sourceUntestedID, sourceAnchoredOnlyID, grandID})
})
if _, err := assets.AddAgentScope(source.ID, "root_domain", sourceTestedDomain, "test", "agent"); err != nil {
t.Fatal(err)
}
if _, err := assets.AddAgentScope(source.ID, "root_domain", sourceUntestedDomain, "test", "agent"); err != nil {
t.Fatal(err)
}
if _, err := assets.AddAgentScope(grand.ID, "root_domain", grandDomain, "test", "agent"); err != nil {
t.Fatal(err)
}
if _, err := d.Exploration(source.ExplorationID).AddNode(KindFact, map[string]any{"summary": "tested source asset"}, 0, "confirmed", "worker", []int64{sourceTestedID}); err != nil {
t.Fatal(err)
}
if _, err := d.Exploration(source.ExplorationID).AddNode(KindFact, map[string]any{"summary": "anchored-only source asset"}, 0, "confirmed", "worker", []int64{sourceAnchoredOnlyID}); err != nil {
t.Fatal(err)
}
if _, err := d.Exploration(grand.ExplorationID).AddNode(KindFact, map[string]any{"summary": "indirect tested asset"}, 0, "confirmed", "worker", []int64{grandID}); err != nil {
t.Fatal(err)
}
scopes, err := assets.ListTaskScopeWithSources(current.ID)
if err != nil {
t.Fatal(err)
}
if len(scopes) != 2 || scopes[0].TaskID != source.ID || scopes[1].TaskID != source.ID {
t.Fatalf("direct source scopes only: %+v", scopes)
}
cov, err := assets.TaskCoverageWithSources(current.ID)
if err != nil {
t.Fatal(err)
}
if cov.Denominator != 3 || cov.Tested != 2 {
t.Fatalf("combined coverage: %+v", cov)
}
untested, total, err := assets.ListUntestedAssetsWithSources(current.ID, "", 10, 0)
if err != nil {
t.Fatal(err)
}
if total != 1 || len(untested) != 1 || untested[0].ID != sourceUntestedID {
t.Fatalf("combined untested assets: total=%d assets=%+v", total, untested)
}
hosts, err := assets.HostsByTaskWithSources(current.ID)
if err != nil {
t.Fatal(err)
}
hostSet := map[string]bool{}
for _, host := range hosts {
hostSet[host] = true
}
if !hostSet[sourceTestedDomain] || !hostSet[sourceUntestedDomain] || !hostSet[sourceAnchoredOnlyDomain] || hostSet[grandDomain] {
t.Fatalf("direct source hosts only: %v", hosts)
}
refs, err := d.Exploration(current.ExplorationID).AssetRefsWithSources(sourceAnchoredOnlyID)
if err != nil || len(refs) != 1 || !refs[0].Inherited || refs[0].SourceTaskID != source.ID || refs[0].Kind != KindFact {
t.Fatalf("source asset refs provenance: refs=%+v err=%v", refs, err)
}
grandRefs, err := d.Exploration(current.ExplorationID).AssetRefsWithSources(grandID)
if err != nil || len(grandRefs) != 0 {
t.Fatalf("indirect source asset refs leaked: refs=%+v err=%v", grandRefs, err)
}
graph, err := assets.BuildCoverageGraph(current.ID, current.ExplorationID)
if err != nil {
t.Fatal(err)
}
seenAnchoredOnly := false
for _, node := range graph.Nodes {
if node.AssetID == sourceAnchoredOnlyID {
seenAnchoredOnly = node.Tested && node.InScope
}
}
if !seenAnchoredOnly {
t.Fatalf("anchored-only source asset missing from coverage graph: %+v", graph.Nodes)
}
}
+296
View File
@@ -0,0 +1,296 @@
package db
import (
"fmt"
"testing"
)
func TestExplorationFlow(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
expID, err := d.CreateExploration("test", "拿下测试目标")
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM explorations WHERE id=$1`, expID) // cascades nodes/edges/activity
es := d.Exploration(expID)
// goal node + two intents
goal, err := es.AddGoal(map[string]any{"text": "getadmin", "vulnclass": "authz"}, "human")
if err != nil {
t.Fatal(err)
}
if _, err := es.AddIntent(map[string]any{"summary": "enumerate endpoints"}, 5, nil, "planner"); err != nil {
t.Fatal(err)
}
i2, err := es.AddIntent(map[string]any{"summary": "test idor"}, 8, nil, "planner")
if err != nil {
t.Fatal(err)
}
// frontier ordered by priority desc → i2(8) before i1(5)
fr, err := es.Frontier(10)
if err != nil {
t.Fatal(err)
}
if len(fr) != 2 || fr[0].ID != i2 {
t.Fatalf("frontier order wrong: %+v", fr)
}
// atomic claim: first wins, second on same id fails
ok, err := es.ClaimIntent(i2, "worker-1")
if err != nil || !ok {
t.Fatalf("claim i2: ok=%v err=%v", ok, err)
}
ok2, _ := es.ClaimIntent(i2, "worker-2")
if ok2 {
t.Fatalf("double-claim should fail")
}
// finding yields from intent, proves goal
find, err := es.AddNode("finding", map[string]any{"vulnclass": "idor", "severity": "high"}, 9, "confirmed", "worker-1", nil)
if err != nil {
t.Fatal(err)
}
if err := es.Link(i2, "yields", find); err != nil {
t.Fatal(err)
}
if err := es.Link(find, "proves", goal); err != nil {
t.Fatal(err)
}
if err := es.SetNodeState(goal, "met"); err != nil {
t.Fatal(err)
}
// lineage: ancestors of the finding traced backward — here {i2, find} joined by
// the yields edge. The proves→goal edge is DOWNSTREAM (goal must be excluded),
// and the unrelated intent i1 is not on any path to the finding (excluded too).
lnNodes, lnEdges, err := es.FindingLineage(find)
if err != nil {
t.Fatal(err)
}
got := map[int64]bool{}
for _, n := range lnNodes {
got[n.ID] = true
}
if len(lnNodes) != 2 || !got[find] || !got[i2] {
t.Fatalf("lineage nodes: want {i2,find}, got %+v", lnNodes)
}
if got[goal] {
t.Fatalf("lineage must exclude the proved goal (it is downstream of the finding)")
}
if len(lnEdges) != 1 || lnEdges[0].From != i2 || lnEdges[0].To != find || lnEdges[0].Rel != "yields" {
t.Fatalf("lineage edges: want i2-yields->find, got %+v", lnEdges)
}
// activity poll by id cursor
id1, err := es.AppendActivity(Activity{Worker: "worker-1", Kind: "tool_use", Tool: "Bash", Summary: "ran curl", Detail: "full output"})
if err != nil {
t.Fatal(err)
}
items, cursor, err := es.ActivityList(nil, 0, 100)
if err != nil {
t.Fatal(err)
}
if len(items) != 1 || cursor != id1 {
t.Fatalf("activity list: items=%d cursor=%d", len(items), cursor)
}
det, _ := es.ActivityDetail(id1)
if det != "full output" {
t.Fatalf("detail want 'full output', got %q", det)
}
// incremental: nothing new after cursor
items2, _, _ := es.ActivityList(nil, cursor, 100)
if len(items2) != 0 {
t.Fatalf("incremental poll should be empty, got %d", len(items2))
}
// stats
st, _ := es.Stats()
if st["intent"] != 2 || st["goal"] != 1 || st["finding"] != 1 {
t.Fatalf("stats: %+v", st)
}
}
func TestIntentPauseResumeAndCancelCleanup(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
expID, err := d.CreateExploration("test", "worker control cleanup")
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM explorations WHERE id=$1`, expID)
es := d.Exploration(expID)
assetID, err := d.Assets().UpsertRootDomain(UpsertRootDomainReq{
Domain: fmt.Sprintf("cancel-intent-%d.invalid", expID),
})
if err != nil {
t.Fatal(err)
}
defer deleteAsset(d, assetID)
otherIntent, err := es.AddIntent(map[string]any{"summary": "keep intent"}, 1, nil, "planner")
if err != nil {
t.Fatal(err)
}
intentID, err := es.AddIntent(map[string]any{"summary": "cancel intent"}, 10, nil, "planner")
if err != nil {
t.Fatal(err)
}
if claimed, err := es.ClaimIntent(intentID, "worker-control-test"); err != nil || !claimed {
t.Fatalf("initial claim: claimed=%v err=%v", claimed, err)
}
if err := es.SetIntentState(intentID, "paused"); err != nil {
t.Fatal(err)
}
frontier, err := es.Frontier(10)
if err != nil {
t.Fatal(err)
}
for _, node := range frontier {
if node.ID == intentID {
t.Fatalf("paused intent %d must not enter frontier", intentID)
}
}
if claimed, err := es.ClaimIntent(intentID, "worker-while-paused"); err != nil || claimed {
t.Fatalf("paused intent claim: claimed=%v err=%v", claimed, err)
}
if err := es.SetIntentState(intentID, "open"); err != nil {
t.Fatal(err)
}
if claimed, err := es.ClaimIntent(intentID, "worker-after-resume"); err != nil || !claimed {
t.Fatalf("resumed intent claim: claimed=%v err=%v", claimed, err)
}
if err := es.SetIntentState(intentID, "paused"); err != nil {
t.Fatal(err)
}
directFact, err := es.AddNode("fact", map[string]any{"summary": "remove fact"}, 0, "confirmed", "worker", []int64{assetID})
if err != nil {
t.Fatal(err)
}
directFinding, err := es.AddNode("finding", map[string]any{"summary": "remove finding"}, 0, "confirmed", "worker", nil)
if err != nil {
t.Fatal(err)
}
keptFact, err := es.AddNode("fact", map[string]any{"summary": "keep fact"}, 0, "confirmed", "other-worker", nil)
if err != nil {
t.Fatal(err)
}
if err := es.Link(intentID, RelYields, directFact); err != nil {
t.Fatal(err)
}
if err := es.Link(intentID, RelYields, directFinding); err != nil {
t.Fatal(err)
}
findingRowID, err := es.AddStandaloneFinding(0, directFinding, "test", "cancelled finding", SeverityHigh, "summary", "evidence", "worker", []int64{assetID})
if err != nil {
t.Fatal(err)
}
activityID, err := es.AppendActivity(Activity{NodeID: &intentID, Worker: "worker", Kind: "result", Summary: "remove activity"})
if err != nil {
t.Fatal(err)
}
keptActivityID, err := es.AppendActivity(Activity{NodeID: &otherIntent, Worker: "other-worker", Kind: "result", Summary: "keep activity"})
if err != nil {
t.Fatal(err)
}
cleanup, err := es.CancelIntent(intentID)
if err != nil {
t.Fatal(err)
}
if cleanup.Intents != 1 || cleanup.Facts != 1 || cleanup.Findings != 1 || cleanup.Activities != 1 {
t.Fatalf("unexpected cleanup counts: %+v", cleanup)
}
for _, nodeID := range []int64{intentID, directFact, directFinding} {
node, err := es.GetNode(nodeID)
if err != nil {
t.Fatal(err)
}
if node != nil {
t.Fatalf("node %d survived intent cancellation", nodeID)
}
}
for _, nodeID := range []int64{otherIntent, keptFact} {
node, err := es.GetNode(nodeID)
if err != nil || node == nil {
t.Fatalf("unrelated node %d removed: node=%v err=%v", nodeID, node, err)
}
}
assertCount := func(query string, want int, args ...any) {
t.Helper()
var got int
if err := d.QueryRow(query, args...).Scan(&got); err != nil {
t.Fatal(err)
}
if got != want {
t.Fatalf("query count=%d, want %d: %s", got, want, query)
}
}
assertCount(`SELECT COUNT(*) FROM findings WHERE id=$1`, 0, findingRowID)
assertCount(`SELECT COUNT(*) FROM activity WHERE id=$1`, 0, activityID)
assertCount(`SELECT COUNT(*) FROM activity WHERE id=$1`, 1, keptActivityID)
assertCount(`SELECT COUNT(*) FROM exploration_edges WHERE exploration_id=$1 AND (src_id=$2 OR dst_id=$2)`, 0, expID, intentID)
assertCount(`SELECT COUNT(*) FROM exploration_anchors WHERE node_id=$1`, 0, directFact)
assertCount(`SELECT COUNT(*) FROM assets WHERE id=$1`, 1, assetID)
}
// TestNodesPageQueryMatchesID verifies the 播报板 search filters on node id (both
// the bare number and the「#id」form the UI shows) in addition to payload/origin.
func TestNodesPageQueryMatchesID(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
expID, err := d.CreateExploration("test", "id 搜索")
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM explorations WHERE id=$1`, expID)
es := d.Exploration(expID)
target, err := es.AddNode("fact", map[string]any{"summary": "needle-alpha"}, 0, "confirmed", "worker-a", nil)
if err != nil {
t.Fatal(err)
}
other, err := es.AddNode("fact", map[string]any{"summary": "unrelated-beta"}, 0, "confirmed", "worker-b", nil)
if err != nil {
t.Fatal(err)
}
onlyTarget := func(label, q string) {
t.Helper()
nodes, total, err := es.NodesPage(NodeFilter{Query: q}, 1, 50)
if err != nil {
t.Fatalf("%s: %v", label, err)
}
if total != 1 || len(nodes) != 1 || nodes[0].ID != target {
t.Fatalf("%s: q=%q total=%d nodes=%+v, want single node %d", label, q, total, nodes, target)
}
}
onlyTarget("bare id", fmt.Sprint(target))
onlyTarget("hash id", "#"+fmt.Sprint(target))
onlyTarget("payload still works", "needle-alpha")
// A non-matching numeric id returns nothing (and does not accidentally match other).
if nodes, total, err := es.NodesPage(NodeFilter{Query: fmt.Sprint(target + other + 1000)}, 1, 50); err != nil {
t.Fatal(err)
} else if total != 0 || len(nodes) != 0 {
t.Fatalf("non-existent id: total=%d nodes=%+v, want empty", total, nodes)
}
}
+130
View File
@@ -0,0 +1,130 @@
package db
import (
"fmt"
"testing"
)
func TestTokenStatsBySessionUsesCompleteIntentHistory(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
expID, err := d.CreateExploration("token sessions", "aggregate complete session history")
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM explorations WHERE id=$1`, expID)
store := d.Exploration(expID)
intentID, err := store.AddIntent(map[string]any{"summary": "worker session"}, 1, nil, "planner")
if err != nil {
t.Fatal(err)
}
factID, err := store.AddNode(KindFact, map[string]any{"summary": "not a worker"}, 0, "confirmed", "worker", nil)
if err != nil {
t.Fatal(err)
}
otherExpID, err := d.CreateExploration("other token sessions", "foreign intent must not match")
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM explorations WHERE id=$1`, otherExpID)
foreignIntentID, err := d.Exploration(otherExpID).AddIntent(map[string]any{"summary": "foreign worker"}, 1, nil, "planner")
if err != nil {
t.Fatal(err)
}
appendResult := func(nodeID *int64, worker string, input, output, read, write int) {
t.Helper()
if _, appendErr := store.AppendActivity(Activity{
NodeID: nodeID,
Worker: worker,
Kind: "result",
InputTokens: &input,
OutputTokens: &output,
CacheReadTokens: &read,
CacheWriteTokens: &write,
}); appendErr != nil {
t.Fatal(appendErr)
}
}
appendResult(nil, "mainagent", 11, 12, 13, 14)
appendResult(nil, "planner", 21, 22, 23, 24)
// More than ActivityPage's default limit proves this aggregate reads the
// persisted history directly instead of summing the currently loaded page.
const completedRuns = 205
for i := 0; i < completedRuns; i++ {
appendResult(&intentID, fmt.Sprintf("work#%d", i%3+1), 1, 2, 3, 4)
}
// All three rows remain part of the legacy worker/whole-task totals, but none
// represents a local Worker session and therefore none may create intent:*.
appendResult(&factID, "work#fact", 31, 32, 33, 34)
appendResult(nil, "work#missing-node", 41, 42, 43, 44)
appendResult(&foreignIntentID, "work#foreign", 51, 52, 53, 54)
ignoredInput := 1000
if _, err := store.AppendActivity(Activity{
NodeID: &intentID,
Worker: "work#1",
Kind: "usage",
InputTokens: &ignoredInput,
}); err != nil {
t.Fatal(err)
}
sessions, err := store.TokenStatsBySession()
if err != nil {
t.Fatal(err)
}
got := make(map[string]SessionTokenUsage, len(sessions))
for _, usage := range sessions {
got[usage.Session] = usage
}
if len(got) != 3 {
t.Fatalf("sessions = %+v, want only main, plan, and the local intent", sessions)
}
assertSessionTokens(t, got["main:0"], "main:0", 11, 12, 13, 14)
assertSessionTokens(t, got["plan"], "plan", 21, 22, 23, 24)
assertSessionTokens(t, got[fmt.Sprintf("intent:%d", intentID)], fmt.Sprintf("intent:%d", intentID),
completedRuns, completedRuns*2, completedRuns*3, completedRuns*4)
workers, err := store.TokenStatsByWorker()
if err != nil {
t.Fatal(err)
}
if len(workers) != 8 { // main, planner, 3 executors, and the 3 deliberately invalid session rows
t.Fatalf("legacy workers changed: got %d entries: %+v", len(workers), workers)
}
total, err := store.TokenTotal()
if err != nil {
t.Fatal(err)
}
assertTokenUsage(t, total,
11+21+completedRuns+31+41+51,
12+22+completedRuns*2+32+42+52,
13+23+completedRuns*3+33+43+53,
14+24+completedRuns*4+34+44+54)
}
func assertSessionTokens(t *testing.T, got SessionTokenUsage, session string, input, output, read, write int) {
t.Helper()
if got.Session != session || got.InputTokens != input || got.OutputTokens != output ||
got.CacheReadTokens != read || got.CacheWriteTokens != write {
t.Fatalf("session %q = %+v, want input=%d output=%d cache-read=%d cache-write=%d",
session, got, input, output, read, write)
}
}
func assertTokenUsage(t *testing.T, got TokenUsage, input, output, read, write int) {
t.Helper()
if got.InputTokens != input || got.OutputTokens != output ||
got.CacheReadTokens != read || got.CacheWriteTokens != write {
t.Fatalf("token total = %+v, want input=%d output=%d cache-read=%d cache-write=%d",
got, input, output, read, write)
}
}
+676
View File
@@ -0,0 +1,676 @@
package db
import (
"encoding/json"
"sort"
"strconv"
"strings"
"time"
)
// ---------------------------------------------------------------------------
// Finding asset tree — the「按资产」view of the global findings list.
//
// 层级与 BuildCoverageGraph 同源(company → root_domain/ip/app → subdomain →
// service → endpoint),但那份是「某个任务的范围内资产」的力导向图,这份是
// 「全库有发现的资产」的树:只收有发现的资产及其祖先链,节点带子树聚合计数。
// 两者的父子优先级规则必须保持一致,改一处请对照 task_scope.go 改另一处。
// ---------------------------------------------------------------------------
// FindingUnassignedAsset 是「未关联资产」的节点 key,也是列表接口的筛选哨兵:
// 命中 asset_ids 为空、或所指资产已被删除的发现。
const FindingUnassignedAsset = "__none__"
// findingAssetTreeMaxNodes 是返回给前端的节点上限。超出时自底向上丢弃整层
// (endpoint 优先,其次 service):它们的计数已经累加进父节点,丢节点不丢数字。
const findingAssetTreeMaxNodes = 3000
// FindingAssetNode 是资产树的一个节点。Key 与覆盖图同构:资产行是 "a:<id>"、
// 企业是 "c:<id>"、没有资产行的根域名是合成的 "r:<domain>"、未关联桶是 "__none__"。
type FindingAssetNode struct {
Key string `json:"key"`
Parent string `json:"parent,omitempty"`
Kind string `json:"kind"` // company|root_domain|subdomain|ip|service|app|endpoint|none
Label string `json:"label"`
AssetID int64 `json:"asset_id,omitempty"`
CompanyID int64 `json:"company_id,omitempty"`
// Self 是直接挂在该资产上的发现数;Total 含全部子孙且按 finding 去重
// (一个发现挂多个资产时,只在共同祖先上计一次)。
Self int `json:"self"`
Total int `json:"total"`
Critical int `json:"critical"`
High int `json:"high"`
Medium int `json:"medium"`
Low int `json:"low"`
LastFoundAt time.Time `json:"last_found_at"`
}
// FindingAssetTree 是整棵树的一次性快照。Nodes 已排好序:同一父节点下按发现数
// 降序、标签升序,「未关联资产」恒在最后。
type FindingAssetTree struct {
Nodes []FindingAssetNode `json:"nodes"`
FindingTotal int `json:"finding_total"`
// Truncated=true 表示为控制体积丢弃了 DroppedKinds 里的层级。
Truncated bool `json:"truncated"`
DroppedKinds []string `json:"dropped_kinds,omitempty"`
}
// assetRow 是构树需要的资产字段子集。
type assetRow struct {
id int64
kind string
companyID int64
domain string
rootDomain string
ip string
url string
port int
serviceType string
appName string
}
func (a *assetRow) coverageNode() CoverageGraphNode {
return CoverageGraphNode{
Kind: a.kind, Domain: a.domain, RootDomain: a.rootDomain, IP: a.ip,
URL: a.url, Port: a.port, ServiceType: a.serviceType, AppName: a.appName,
}
}
// label 复用覆盖图的标签规则(URL > domain > ip > app_name > root_domain),但没有
// URL 的服务(SMB、非 HTTP 端口等)要补上端口:否则它的标签会和宿主 IP/域名那行
// 一模一样,树上父子两行看起来完全重复。
func (a *assetRow) label() string {
if a.kind == "service" && a.url == "" {
if host, port := a.hostPort(); host != "" && port > 0 {
return host + ":" + strconv.Itoa(port)
}
}
n := a.coverageNode()
n.Key = assetKey(a.id)
return coverageNodeLabel(&n)
}
// hostPort 与覆盖图一致:优先 domain,其次 URL 里的 host,最后 ip。
func (a *assetRow) hostPort() (string, int) {
n := a.coverageNode()
return hostPortOf(&n)
}
const findingAssetSelectCols = `a.id, a.type, COALESCE(a.company_id,0),
COALESCE(a.domain,''), COALESCE(a.root_domain,''), COALESCE(a.ip,''),
COALESCE(a.url,''), COALESCE(a.port,0), COALESCE(a.service_type,''),
COALESCE(a.app_name,'')`
func scanAssetRows(rows interface {
Next() bool
Scan(...any) error
Err() error
Close() error
}) ([]*assetRow, error) {
defer rows.Close()
var out []*assetRow
for rows.Next() {
a := &assetRow{}
if err := rows.Scan(&a.id, &a.kind, &a.companyID, &a.domain, &a.rootDomain,
&a.ip, &a.url, &a.port, &a.serviceType, &a.appName); err != nil {
return nil, err
}
out = append(out, a)
}
return out, rows.Err()
}
// findingAssetHit 是一条发现在构树阶段需要的最小信息。
type findingAssetHit struct {
severity string
ts time.Time
assetIDs []int64
}
// BuildFindingAssetTree 按当前筛选构建资产树。AssetScope 自身不参与(否则树会随
// 选中节点塌缩成一条链)。
func (d *DB) BuildFindingAssetTree(f FindingFilter) (*FindingAssetTree, error) {
return d.buildFindingAssetTree(f, findingAssetTreeMaxNodes)
}
// buildFindingAssetTree 是带节点上限的内部实现。maxNodes<=0 表示不截断——解析
// AssetScope 时必须用这个模式,否则被丢掉的 endpoint 会让子树 id 集合不全。
func (d *DB) buildFindingAssetTree(f FindingFilter, maxNodes int) (*FindingAssetTree, error) {
f.AssetScope = ""
f.assetIDs, f.assetNone, f.assetMiss = nil, false, false
where, args := f.where()
rows, err := d.Query(`SELECT COALESCE(f.severity,''), f.created_at,
COALESCE(f.asset_ids::text,'[]')
FROM findings f LEFT JOIN tasks t ON f.task_id = t.id`+where, args...)
if err != nil {
return nil, err
}
hits := []findingAssetHit{}
assetIDs := map[int64]bool{}
for rows.Next() {
var h findingAssetHit
var aidsJSON string
if err := rows.Scan(&h.severity, &h.ts, &aidsJSON); err != nil {
rows.Close()
return nil, err
}
_ = json.Unmarshal([]byte(aidsJSON), &h.assetIDs)
for _, id := range h.assetIDs {
if id > 0 {
assetIDs[id] = true
}
}
hits = append(hits, h)
}
if err := rows.Err(); err != nil {
rows.Close()
return nil, err
}
rows.Close()
tree := &FindingAssetTree{Nodes: []FindingAssetNode{}, FindingTotal: len(hits)}
byID, err := d.loadFindingAssetRows(assetIDs)
if err != nil {
return nil, err
}
nodes, parentOf := d.assembleFindingAssetNodes(byID)
if err := d.attachCompanyNodes(nodes, parentOf); err != nil {
return nil, err
}
// 计数:一条发现沿它每个资产的祖先链向上,收集去重后的 key 集合再逐个 +1,
// 所以父节点不会因为一条发现挂了多个子资产而重复计数。
unassigned := &FindingAssetNode{Key: FindingUnassignedAsset, Kind: "none", Label: "未关联资产"}
touched := map[string]bool{}
for _, h := range hits {
clear(touched)
var direct []*FindingAssetNode
for _, id := range h.assetIDs {
node := nodes[assetKey(id)]
if node == nil {
continue
}
direct = append(direct, node)
for key := node.Key; key != ""; key = parentOf[key] {
touched[key] = true
}
}
if len(direct) == 0 {
countFinding(unassigned, h)
unassigned.Self++
continue
}
for _, node := range direct {
node.Self++
}
for key := range touched {
countFinding(nodes[key], h)
}
}
for _, node := range nodes {
if node.Total > 0 {
tree.Nodes = append(tree.Nodes, *node)
}
}
if unassigned.Total > 0 {
tree.Nodes = append(tree.Nodes, *unassigned)
}
sortFindingAssetNodes(tree.Nodes)
truncateFindingAssetTree(tree, maxNodes)
return tree, nil
}
// countFinding 把一条发现累加到节点上(总数 / 严重度分桶 / 最近发现时间)。
func countFinding(n *FindingAssetNode, h findingAssetHit) {
if n == nil {
return
}
n.Total++
switch h.severity {
case "critical":
n.Critical++
case "high":
n.High++
case "medium":
n.Medium++
case "low":
n.Low++
}
if h.ts.After(n.LastFoundAt) {
n.LastFoundAt = h.ts
}
}
// loadFindingAssetRows 读取命中的资产行,并逐轮补齐祖先(service 的宿主域名/IP、
// 子域名的根域名)。祖先自身可能没有任何发现,但树需要它们才能成形。
func (d *DB) loadFindingAssetRows(ids map[int64]bool) (map[int64]*assetRow, error) {
byID := map[int64]*assetRow{}
if len(ids) == 0 {
return byID, nil
}
idList := make([]int64, 0, len(ids))
for id := range ids {
idList = append(idList, id)
}
rows, err := d.Query(`SELECT `+findingAssetSelectCols+` FROM assets a WHERE a.id = ANY($1::bigint[])`, idList)
if err != nil {
return nil, err
}
found, err := scanAssetRows(rows)
if err != nil {
return nil, err
}
for _, a := range found {
byID[a.id] = a
}
// 每轮找出还缺父节点的宿主标识,批量补一层;层数固定(endpoint→service→
// subdomain/ip→root_domain),4 轮足够收敛。
for range 4 {
want := missingParents(byID)
if want.empty() {
break
}
added, err := d.loadAssetsByHost(want, byID)
if err != nil {
return nil, err
}
if added == 0 {
break
}
}
return byID, nil
}
// missingHosts 是一轮补齐里要去库里找的宿主标识,按目标资产类型分开。
type missingHosts struct {
services []string // endpoint 的宿主(找 service 行)
domains []string // service/endpoint 的宿主域名(找 subdomain 行)
ips []string // service/endpoint 的宿主 IP(找 ip 行)
roots []string // 子域名的根域名(找 root_domain 行)
}
func (m missingHosts) empty() bool {
return len(m.services) == 0 && len(m.domains) == 0 && len(m.ips) == 0 && len(m.roots) == 0
}
// missingParents 汇总还没被加载的宿主:service(供 endpoint 挂靠)、子域名/IP(供
// service 与 endpoint 挂靠)与根域名(供子域名挂靠)。
func missingParents(byID map[int64]*assetRow) missingHosts {
haveService := map[string]bool{}
haveDomain := map[string]bool{}
haveIP := map[string]bool{}
haveRoot := map[string]bool{}
for _, a := range byID {
switch a.kind {
case "service":
if host, _ := a.hostPort(); host != "" {
haveService[host] = true
}
case "subdomain":
haveDomain[a.domain] = true
case "ip":
haveIP[a.ip] = true
case "root_domain":
haveRoot[a.domain] = true
}
}
wantService := map[string]bool{}
wantDomain := map[string]bool{}
wantIP := map[string]bool{}
wantRoot := map[string]bool{}
for _, a := range byID {
switch a.kind {
case "service", "endpoint":
host, _ := a.hostPort()
// endpoint 先找同宿主的 service;端口对不上的 service 会以 Total=0
// 被最终过滤掉,不会污染树。
if a.kind == "endpoint" && host != "" && !haveService[host] {
wantService[host] = true
}
if host != "" && !haveDomain[host] && !haveIP[host] && !haveRoot[host] {
if isIPLiteral(host) {
wantIP[host] = true
} else {
wantDomain[host] = true
}
}
if a.ip != "" && !haveIP[a.ip] {
wantIP[a.ip] = true
}
case "subdomain":
if a.rootDomain != "" && !haveRoot[a.rootDomain] {
wantRoot[a.rootDomain] = true
}
}
}
return missingHosts{
services: keysOf(wantService),
domains: keysOf(wantDomain),
ips: keysOf(wantIP),
roots: keysOf(wantRoot),
}
}
func keysOf(m map[string]bool) []string {
if len(m) == 0 {
return nil
}
out := make([]string, 0, len(m))
for k := range m {
out = append(out, k)
}
return out
}
// isIPLiteral 粗判一个 host 是不是 IP 字面量(用于决定去 ip 还是 subdomain 表找宿主)。
func isIPLiteral(host string) bool {
if strings.Contains(host, ":") {
return true // IPv6
}
if host == "" {
return false
}
for _, part := range strings.Split(host, ".") {
if part == "" {
return false
}
if _, err := strconv.Atoi(part); err != nil {
return false
}
}
return strings.Count(host, ".") == 3
}
// loadAssetsByHost 按宿主标识批量补齐资产行,返回本轮新增的行数。
func (d *DB) loadAssetsByHost(want missingHosts, byID map[int64]*assetRow) (int, error) {
added := 0
load := func(q string, arg []string) error {
if len(arg) == 0 {
return nil
}
rows, err := d.Query(q, arg)
if err != nil {
return err
}
found, err := scanAssetRows(rows)
if err != nil {
return err
}
for _, a := range found {
if _, ok := byID[a.id]; ok {
continue
}
byID[a.id] = a
added++
}
return nil
}
if err := load(`SELECT `+findingAssetSelectCols+` FROM assets a
WHERE a.type='service' AND (a.domain = ANY($1::text[]) OR a.ip = ANY($1::text[]))`, want.services); err != nil {
return 0, err
}
if err := load(`SELECT `+findingAssetSelectCols+` FROM assets a
WHERE a.type='subdomain' AND a.domain = ANY($1::text[])`, want.domains); err != nil {
return 0, err
}
if err := load(`SELECT `+findingAssetSelectCols+` FROM assets a
WHERE a.type='ip' AND a.ip = ANY($1::text[])`, want.ips); err != nil {
return 0, err
}
if err := load(`SELECT `+findingAssetSelectCols+` FROM assets a
WHERE a.type='root_domain' AND a.domain = ANY($1::text[])`, want.roots); err != nil {
return 0, err
}
return added, nil
}
// assembleFindingAssetNodes 把资产行变成节点并连上父子关系。父节点缺位时(库里
// 根本没有那条根域名资产)合成 "r:<domain>" 占位节点,与覆盖图的处理一致。
func (d *DB) assembleFindingAssetNodes(byID map[int64]*assetRow) (map[string]*FindingAssetNode, map[string]string) {
nodes := map[string]*FindingAssetNode{}
parentOf := map[string]string{}
rootByDomain := map[string]string{}
subByDomain := map[string]string{}
ipByAddr := map[string]string{}
svcByHost := map[string]string{}
svcByHostPort := map[string]string{}
for _, a := range byID {
key := assetKey(a.id)
nodes[key] = &FindingAssetNode{
Key: key, Kind: a.kind, Label: a.label(),
AssetID: a.id, CompanyID: a.companyID,
}
switch a.kind {
case "root_domain":
if a.domain != "" {
rootByDomain[a.domain] = key
}
case "subdomain":
if a.domain != "" {
subByDomain[a.domain] = key
}
case "ip":
if a.ip != "" {
ipByAddr[a.ip] = key
}
case "service":
if host, port := a.hostPort(); host != "" {
svcByHost[host] = key
svcByHostPort[host+"|"+strconv.Itoa(port)] = key
}
}
}
// 子域名的根域名在库里没有资产行时,合成一个占位根,免得子域名散成顶层。
for _, a := range byID {
if a.kind != "subdomain" || a.rootDomain == "" {
continue
}
if _, ok := rootByDomain[a.rootDomain]; ok {
continue
}
key := "r:" + a.rootDomain
nodes[key] = &FindingAssetNode{Key: key, Kind: "root_domain", Label: a.rootDomain}
rootByDomain[a.rootDomain] = key
}
firstOf := func(keys ...string) string {
for _, k := range keys {
if k != "" {
if _, ok := nodes[k]; ok {
return k
}
}
}
return ""
}
for _, a := range byID {
key := assetKey(a.id)
var parent string
switch a.kind {
case "subdomain":
parent = firstOf(rootByDomain[a.rootDomain])
case "service":
host, _ := a.hostPort()
parent = firstOf(subByDomain[a.domain], subByDomain[host],
ipByAddr[a.ip], ipByAddr[host], rootByDomain[a.rootDomain], rootByDomain[host])
case "endpoint":
host, port := a.hostPort()
parent = firstOf(svcByHostPort[host+"|"+strconv.Itoa(port)], svcByHost[host],
subByDomain[host], subByDomain[a.domain], ipByAddr[host], ipByAddr[a.ip],
rootByDomain[a.rootDomain], rootByDomain[host])
}
if parent != "" && parent != key {
parentOf[key] = parent
nodes[key].Parent = parent
}
}
return nodes, parentOf
}
// attachCompanyNodes 给顶层资产(根域名 / IP / 应用)补企业父节点——只有资产确实
// 归属了企业才会出现企业层,没归属的资产仍然自己就是顶层。
func (d *DB) attachCompanyNodes(nodes map[string]*FindingAssetNode, parentOf map[string]string) error {
want := map[int64]bool{}
for _, n := range nodes {
if n.Parent != "" || n.CompanyID <= 0 {
continue
}
switch n.Kind {
case "root_domain", "ip", "app":
want[n.CompanyID] = true
}
}
if len(want) == 0 {
return nil
}
ids := make([]int64, 0, len(want))
for id := range want {
ids = append(ids, id)
}
rows, err := d.Query(`SELECT id, COALESCE(name,'') FROM companies WHERE id = ANY($1::bigint[])`, ids)
if err != nil {
return err
}
defer rows.Close()
names := map[int64]string{}
for rows.Next() {
var id int64
var name string
if err := rows.Scan(&id, &name); err != nil {
return err
}
names[id] = name
}
if err := rows.Err(); err != nil {
return err
}
for id, name := range names {
key := companyKey(id)
if _, ok := nodes[key]; ok {
continue
}
if name == "" {
name = "企业 #" + strconv.FormatInt(id, 10)
}
nodes[key] = &FindingAssetNode{Key: key, Kind: "company", Label: name, CompanyID: id}
}
for _, n := range nodes {
if n.Parent != "" || n.CompanyID <= 0 || n.Kind == "company" {
continue
}
switch n.Kind {
case "root_domain", "ip", "app":
key := companyKey(n.CompanyID)
if _, ok := nodes[key]; !ok {
continue
}
n.Parent = key
parentOf[n.Key] = key
}
}
return nil
}
// sortFindingAssetNodes 排序:发现多的在前,同数按标签;「未关联资产」恒在最后。
// 前端按数组顺序挂子节点,所以只要同一父节点下的相对顺序正确即可。
func sortFindingAssetNodes(nodes []FindingAssetNode) {
sort.SliceStable(nodes, func(i, j int) bool {
a, b := nodes[i], nodes[j]
if (a.Kind == "none") != (b.Kind == "none") {
return b.Kind == "none"
}
if a.Total != b.Total {
return a.Total > b.Total
}
return a.Label < b.Label
})
}
// truncateFindingAssetTree 在节点过多时整层丢弃(先 endpoint 再 service)。计数已
// 累加到父节点,丢的只是可展开的细节层级。
func truncateFindingAssetTree(tree *FindingAssetTree, maxNodes int) {
if maxNodes <= 0 || len(tree.Nodes) <= maxNodes {
return
}
for _, kind := range []string{"endpoint", "service"} {
kept := tree.Nodes[:0]
for _, n := range tree.Nodes {
if n.Kind == kind {
continue
}
kept = append(kept, n)
}
tree.Nodes = kept
tree.Truncated = true
tree.DroppedKinds = append(tree.DroppedKinds, kind)
if len(tree.Nodes) <= maxNodes {
return
}
}
}
// applyAssetScope 把 AssetScope(节点 key)解析成可用于 SQL 的资产 id 集合。选中
// 一个节点等于选中它的整棵子树,所以要先把树建出来再收集子孙。
func (d *DB) applyAssetScope(f FindingFilter) (FindingFilter, error) {
scope := strings.TrimSpace(f.AssetScope)
f.assetIDs, f.assetNone, f.assetMiss = nil, false, false
if scope == "" {
return f, nil
}
if scope == FindingUnassignedAsset {
f.assetNone = true
return f, nil
}
// 不截断:被丢掉的 endpoint 同样要参与 id 收集,否则列表会少数据。
tree, err := d.buildFindingAssetTree(f, 0)
if err != nil {
return f, err
}
children := map[string][]FindingAssetNode{}
byKey := map[string]FindingAssetNode{}
for _, n := range tree.Nodes {
byKey[n.Key] = n
children[n.Parent] = append(children[n.Parent], n)
}
if _, ok := byKey[scope]; !ok {
// 选中的节点在当前筛选下已经不存在,结果应当为空而不是退化成不过滤。
f.assetMiss = true
return f, nil
}
seen := map[string]bool{scope: true}
queue := []string{scope}
for len(queue) > 0 {
key := queue[0]
queue = queue[1:]
if id := byKey[key].AssetID; id > 0 {
f.assetIDs = append(f.assetIDs, id)
}
for _, child := range children[key] {
if seen[child.Key] {
continue
}
seen[child.Key] = true
queue = append(queue, child.Key)
}
}
if len(f.assetIDs) == 0 {
f.assetMiss = true
}
return f, nil
}
// assetIDContainments 把资产 id 变成 jsonb 包含判断的右操作数集合,配合
// idx_findings_asset_ids(GIN jsonb_path_ops)使用。
func assetIDContainments(ids []int64) []string {
out := make([]string, 0, len(ids))
for _, id := range ids {
out = append(out, "["+strconv.FormatInt(id, 10)+"]")
}
return out
}
+233
View File
@@ -0,0 +1,233 @@
package db
import (
"strconv"
"strings"
"testing"
)
// cleanupTreeFixtures 删除一个用例造出来的资产与发现。必须用 defer 注册(而不是
// t.Cleanup):t.Cleanup 跑在测试函数返回之后,那时 defer d.Close() 已经把连接关了,
// 清理会静默失败并把脏数据留在共享开发库里。
func cleanupTreeFixtures(d *DB, taskID int64, rootDomains ...string) {
d.Exec(`DELETE FROM assets WHERE root_domain = ANY($1::text[])`, rootDomains) //nolint:errcheck
d.DeleteFindingsByTask(taskID) //nolint:errcheck
}
// seedTreeAsset inserts one asset row.
func seedTreeAsset(t *testing.T, d *DB, kind string, cols map[string]any) int64 {
t.Helper()
names := []string{"type"}
values := []any{kind}
placeholders := []string{"$1"}
for k, v := range cols {
values = append(values, v)
names = append(names, k)
placeholders = append(placeholders, "$"+strconv.Itoa(len(values)))
}
q := "INSERT INTO assets(" + strings.Join(names, ",") + ") VALUES (" +
strings.Join(placeholders, ",") + ") RETURNING id"
var id int64
if err := d.QueryRow(q, values...).Scan(&id); err != nil {
t.Fatalf("seed %s asset: %v", kind, err)
}
return id
}
func nodeByKey(tree *FindingAssetTree, key string) *FindingAssetNode {
for i := range tree.Nodes {
if tree.Nodes[i].Key == key {
return &tree.Nodes[i]
}
}
return nil
}
// TestBuildFindingAssetTree covers the whole shape of the「按资产」tree: the
// root→subdomain→service→endpoint chain gets rebuilt from a finding that only
// points at the leaf, ancestors aggregate their subtree, assets without any
// finding stay out, and a finding whose asset row is gone lands in the
// unassigned bucket.
func TestBuildFindingAssetTree(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
tk, err := d.CreateTask("资产树测试", "目标", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
defer d.DeleteTask(tk.ID)
const root = "tree-test.example"
const sub = "api.tree-test.example"
defer cleanupTreeFixtures(d, tk.ID, root)
rootID := seedTreeAsset(t, d, "root_domain", map[string]any{"domain": root, "root_domain": root})
subID := seedTreeAsset(t, d, "subdomain", map[string]any{"domain": sub, "root_domain": root})
svcID := seedTreeAsset(t, d, "service", map[string]any{
"domain": sub, "root_domain": root, "url": "https://" + sub, "port": 443, "service_type": "http",
})
epID := seedTreeAsset(t, d, "endpoint", map[string]any{
"domain": sub, "root_domain": root, "url": "https://" + sub + "/admin", "port": 443, "method": "GET",
})
// 同域名下另一个服务,不挂任何发现 —— 不应出现在树里。
seedTreeAsset(t, d, "service", map[string]any{
"domain": sub, "root_domain": root, "url": "http://" + sub + ":8080", "port": 8080, "service_type": "http",
})
// 只把发现挂在最深的 endpoint 上,祖先链要靠构树自己补出来。
if _, err := d.AddFinding(tk.ID, 0, "XSS", "反射型 XSS", "high", "s", "e", "w", []int64{epID}); err != nil {
t.Fatal(err)
}
// 直接挂在服务上的一条,用来验证 Self 与 Total 的区别。
if _, err := d.AddFinding(tk.ID, 0, "Info", "信息泄露", "low", "s", "e", "w", []int64{svcID}); err != nil {
t.Fatal(err)
}
// 资产行不存在(已删除资产)→ 未关联桶。
if _, err := d.AddFinding(tk.ID, 0, "Misc", "孤儿", "medium", "s", "e", "w", []int64{999000111}); err != nil {
t.Fatal(err)
}
tree, err := d.BuildFindingAssetTree(FindingFilter{TaskID: strconv.FormatInt(tk.ID, 10)})
if err != nil {
t.Fatal(err)
}
if tree.FindingTotal != 3 {
t.Fatalf("finding_total: want 3, got %d", tree.FindingTotal)
}
rootNode := nodeByKey(tree, assetKey(rootID))
subNode := nodeByKey(tree, assetKey(subID))
svcNode := nodeByKey(tree, assetKey(svcID))
epNode := nodeByKey(tree, assetKey(epID))
for name, n := range map[string]*FindingAssetNode{
"root": rootNode, "subdomain": subNode, "service": svcNode, "endpoint": epNode,
} {
if n == nil {
t.Fatalf("%s node missing from tree", name)
}
}
// 父子链:endpoint → service → subdomain → root_domain。
if epNode.Parent != svcNode.Key {
t.Errorf("endpoint parent: want %s, got %s", svcNode.Key, epNode.Parent)
}
if svcNode.Parent != subNode.Key {
t.Errorf("service parent: want %s, got %s", subNode.Key, svcNode.Parent)
}
if subNode.Parent != rootNode.Key {
t.Errorf("subdomain parent: want %s, got %s", rootNode.Key, subNode.Parent)
}
if rootNode.Parent != "" {
t.Errorf("root parent: want top level, got %s", rootNode.Parent)
}
// 聚合:根域名两条(endpoint 的 high + service 的 low),service 自身一条、子树两条。
if rootNode.Total != 2 || rootNode.High != 1 || rootNode.Low != 1 {
t.Errorf("root totals: want 2/high1/low1, got %d/high%d/low%d", rootNode.Total, rootNode.High, rootNode.Low)
}
if rootNode.Self != 0 {
t.Errorf("root self: want 0 (只是祖先), got %d", rootNode.Self)
}
if svcNode.Total != 2 || svcNode.Self != 1 {
t.Errorf("service total/self: want 2/1, got %d/%d", svcNode.Total, svcNode.Self)
}
if epNode.Total != 1 || epNode.Self != 1 {
t.Errorf("endpoint total/self: want 1/1, got %d/%d", epNode.Total, epNode.Self)
}
// 没有发现的兄弟服务不进树。
for _, n := range tree.Nodes {
if n.Label == "http://"+sub+":8080" {
t.Errorf("asset without findings should be hidden: %+v", n)
}
}
// 未关联桶收下那条指向已删资产的发现。
none := nodeByKey(tree, FindingUnassignedAsset)
if none == nil || none.Total != 1 || none.Medium != 1 {
t.Fatalf("unassigned bucket: want 1 medium, got %+v", none)
}
}
// TestFindingAssetScopeFilter verifies选中一个节点 narrows the findings list to
// that node's whole subtree, and that the unassigned sentinel works too.
func TestFindingAssetScopeFilter(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
tk, err := d.CreateTask("资产筛选测试", "目标", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
defer d.DeleteTask(tk.ID)
const root = "scope-test.example"
const sub = "api.scope-test.example"
const other = "other-scope-test.example"
defer cleanupTreeFixtures(d, tk.ID, root, other)
rootID := seedTreeAsset(t, d, "root_domain", map[string]any{"domain": root, "root_domain": root})
subID := seedTreeAsset(t, d, "subdomain", map[string]any{"domain": sub, "root_domain": root})
otherID := seedTreeAsset(t, d, "root_domain", map[string]any{"domain": other, "root_domain": other})
if _, err := d.AddFinding(tk.ID, 0, "A", "子域名上的", "high", "s", "e", "w", []int64{subID}); err != nil {
t.Fatal(err)
}
if _, err := d.AddFinding(tk.ID, 0, "B", "别的根域名上的", "high", "s", "e", "w", []int64{otherID}); err != nil {
t.Fatal(err)
}
if _, err := d.AddFinding(tk.ID, 0, "C", "没有资产的", "high", "s", "e", "w", nil); err != nil {
t.Fatal(err)
}
// 指向已删除资产的发现,和 asset_ids 为空的一样属于「未关联」——树的桶收下它,
// 列表筛选也必须查得出来,两处口径不一致会让桶上的数字大于点开后的条数。
if _, err := d.AddFinding(tk.ID, 0, "D", "资产已删除", "high", "s", "e", "w", []int64{999000333}); err != nil {
t.Fatal(err)
}
base := FindingFilter{TaskID: strconv.FormatInt(tk.ID, 10)}
cases := []struct {
name string
scope string
want int
}{
{"整棵子树", assetKey(rootID), 1}, // 根域名下只有子域名那条
{"叶子节点", assetKey(subID), 1}, // 子域名自身
{"另一棵树", assetKey(otherID), 1}, // 互不串味
{"未关联", FindingUnassignedAsset, 2}, // asset_ids 为空的 + 指向已删资产的
{"不存在的节点", "a:999000222", 0}, // 当前筛选下没有该节点 → 空结果,不是不过滤
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
f := base
f.AssetScope = tc.scope
items, total, err := d.ListFindingsPage(f, 1, 50)
if err != nil {
t.Fatal(err)
}
if total != tc.want || len(items) != tc.want {
t.Fatalf("scope %s: want %d findings, got total=%d items=%d", tc.scope, tc.want, total, len(items))
}
})
}
// 不带 scope 时四条都在。
if _, total, err := d.ListFindingsPage(base, 1, 50); err != nil || total != 4 {
t.Fatalf("unscoped: want 4, got %d (%v)", total, err)
}
// 树上未关联桶的计数必须与点开后查到的条数一致 —— 这正是两处口径分家时会崩的断言。
tree, err := d.BuildFindingAssetTree(base)
if err != nil {
t.Fatal(err)
}
none := nodeByKey(tree, FindingUnassignedAsset)
if none == nil || none.Total != 2 {
t.Fatalf("unassigned bucket count: want 2, got %+v", none)
}
}
+270
View File
@@ -0,0 +1,270 @@
package db
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"strings"
"time"
)
const FindingRetestAgentKey = "retester"
// 재검증(finding_retest) 종결 사유. finding_retests.error 컬럼에 저장돼 재검증 패널
// (finding-retest-panel) 의 item.error 로 노출된다(사용자 노출, server/conversations.go 의
// 형제 사유 convRetest* 와 같은 컬럼·패널이라 함께 한국어로 둔다 — F9). retestNoConclusionReason
// 은 에이전트가 결론 없이 완료했을 때 caller 가 넘긴 사유를 덮어쓰는 폴백이고,
// retestServiceRestartReason 은 재시작 복구(RecoverFindingRetests)가 미완 재검증을 봉인할 때 쓴다.
// 작은따옴표 없는 상수라 SQL 리터럴 자리에 그대로 이어 붙여도 안전하다.
const (
retestNoConclusionReason = "에이전트가 재검증 결론을 저장하지 않았습니다. 대화를 확인한 뒤 다시 재검증해 주세요"
retestServiceRestartReason = "서비스가 재시작되어 재검증이 중단되었습니다. 다시 시작해 주세요"
)
var ErrRetestNotRunning = errors.New("本次复测已结束或尚未开始,请从漏洞详情发起新的复测")
// FindingRetest is an immutable historical test once its conversation turn ends.
// Snapshot is only loaded for the agent, never sent with the history list.
type FindingRetest struct {
ID int64 `json:"id"`
FindingID int64 `json:"finding_id"`
ConversationID *int64 `json:"conversation_id"`
Status string `json:"status"`
Verdict string `json:"verdict"`
Notes string `json:"notes"`
Snapshot json.RawMessage `json:"snapshot,omitempty"`
Summary string `json:"summary"`
Evidence string `json:"evidence"`
Error string `json:"error"`
CreatedAt time.Time `json:"created_at"`
StartedAt *time.Time `json:"started_at"`
FinishedAt *time.Time `json:"finished_at"`
}
const retestCols = `id, finding_id, conversation_id, status, verdict, notes, summary, evidence, error, created_at, started_at, finished_at`
// ActiveFindingRetest is the small status payload polled by the findings list.
// Finding IDs use the same string representation as the findings API.
type ActiveFindingRetest struct {
ID int64 `json:"id"`
FindingID int64 `json:"finding_id,string"`
ConversationID int64 `json:"conversation_id"`
Status string `json:"status"`
}
func (d *DB) ListActiveFindingRetests(ctx context.Context) ([]ActiveFindingRetest, error) {
rows, err := d.QueryContext(ctx, `SELECT id, finding_id, conversation_id, status FROM finding_retests
WHERE status IN ('pending','running') AND conversation_id IS NOT NULL ORDER BY id`)
if err != nil {
return nil, err
}
defer rows.Close()
items := []ActiveFindingRetest{}
for rows.Next() {
var item ActiveFindingRetest
if err := rows.Scan(&item.ID, &item.FindingID, &item.ConversationID, &item.Status); err != nil {
return nil, err
}
items = append(items, item)
}
return items, rows.Err()
}
func scanRetest(row interface{ Scan(...any) error }) (*FindingRetest, error) {
r := &FindingRetest{}
err := row.Scan(&r.ID, &r.FindingID, &r.ConversationID, &r.Status, &r.Verdict, &r.Notes,
&r.Summary, &r.Evidence, &r.Error, &r.CreatedAt, &r.StartedAt, &r.FinishedAt)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
return r, err
}
// CreateFindingRetest atomically snapshots the source, creates its conversation
// and persists the first message. A finding row lock deduplicates simultaneous
// clicks across clients; an existing active run is returned without dispatching.
func (d *DB) CreateFindingRetest(ctx context.Context, findingID int64, notes string) (*FindingRetest, *Conversation, bool, error) {
tx, err := d.BeginTx(ctx, nil)
if err != nil {
return nil, nil, false, err
}
defer tx.Rollback()
var title string
var snapshot []byte
err = tx.QueryRowContext(ctx, `SELECT COALESCE(NULLIF(f.name,''), NULLIF(f.vulnclass,''), '未分类'),
jsonb_build_object('finding', to_jsonb(f),
'assets', COALESCE((SELECT jsonb_agg(to_jsonb(a)) FROM assets a WHERE f.asset_ids @> to_jsonb(ARRAY[a.id])), '[]'::jsonb),
'constraints', COALESCE((SELECT jsonb_agg(to_jsonb(c)) FROM task_constraints c JOIN tasks t ON t.exploration_id=c.exploration_id WHERE t.id=f.task_id), '[]'::jsonb))
FROM findings f WHERE f.id=$1 FOR UPDATE OF f`, findingID).Scan(&title, &snapshot)
if err != nil {
return nil, nil, false, err
}
r, err := scanRetest(tx.QueryRowContext(ctx, `SELECT `+retestCols+` FROM finding_retests WHERE finding_id=$1 AND status IN ('pending','running')`, findingID))
if err != nil {
return nil, nil, false, err
}
if r != nil {
return r, nil, false, nil
}
// Keep the title within the same limit as ordinary conversations.
if runes := []rune(title); len(runes) > 100 {
title = string(runes[:100])
}
c, err := scanConv(tx.QueryRowContext(ctx, `INSERT INTO conversations(agent_key,title) VALUES ($1,$2) RETURNING `+convCols,
FindingRetestAgentKey, fmt.Sprintf("复测 #%d · %s", findingID, title)))
if err != nil {
return nil, nil, false, err
}
r, err = scanRetest(tx.QueryRowContext(ctx, `INSERT INTO finding_retests(finding_id,conversation_id,notes,snapshot) VALUES ($1,$2,$3,$4) RETURNING `+retestCols,
findingID, c.ID, strings.TrimSpace(notes), snapshot))
if err != nil {
return nil, nil, false, err
}
msg := r.InitialMessage()
_, err = tx.ExecContext(ctx, `INSERT INTO conversation_activities(conversation_id,worker,kind,summary,detail) VALUES ($1,$2,'user',$3,$4)`,
c.ID, FindingRetestAgentKey, fmt.Sprintf("请复测漏洞 #%d", findingID), msg)
if err != nil {
return nil, nil, false, err
}
if err = tx.Commit(); err != nil {
return nil, nil, false, err
}
return r, &c, true, nil
}
func (r *FindingRetest) InitialMessage() string {
msg := fmt.Sprintf("请复测漏洞 #%d。先调用 get_finding_retest_context 读取本会话关联的原始证据与约束,再执行针对性验证,最后调用 record_finding_retest_result 保存结论。", r.FindingID)
if r.Notes != "" {
msg += "\n\n本次复测补充说明:\n" + r.Notes
}
return msg
}
func (d *DB) ListFindingRetests(findingID int64) ([]*FindingRetest, error) {
rows, err := d.Query(`SELECT `+retestCols+` FROM finding_retests WHERE finding_id=$1 ORDER BY id DESC`, findingID)
if err != nil {
return nil, err
}
defer rows.Close()
out := []*FindingRetest{}
for rows.Next() {
r, err := scanRetest(rows)
if err != nil {
return nil, err
}
out = append(out, r)
}
return out, rows.Err()
}
func (d *DB) FindingRetestForConversation(ctx context.Context, conversationID int64) (*FindingRetest, error) {
r, err := scanRetest(d.QueryRowContext(ctx, `SELECT `+retestCols+` FROM finding_retests WHERE conversation_id=$1`, conversationID))
if err != nil || r == nil {
return r, err
}
err = d.QueryRowContext(ctx, `SELECT snapshot FROM finding_retests WHERE id=$1`, r.ID).Scan(&r.Snapshot)
return r, err
}
// FailPendingRetestForConversation seals a conversation's unfinished retest when
// the runner could not even load it — the retest ID is unknown on that path, so
// the conversation ID is the only handle. Without it a transient read error
// leaves the row 'pending' forever: the findings list keeps showing 复测中 and
// every later 发起复测 is deduped against a run that is not happening, with only
// a process restart (RecoverFindingRetests) able to clear it.
func (d *DB) FailPendingRetestForConversation(conversationID int64, reason string) error {
_, err := d.Exec(`UPDATE finding_retests SET status='failed', error=$2, finished_at=now()
WHERE conversation_id=$1 AND status IN ('pending','running')`, conversationID, reason)
return err
}
func (d *DB) StartFindingRetest(ctx context.Context, id int64) (bool, error) {
res, err := d.ExecContext(ctx, `UPDATE finding_retests SET status='running', started_at=now() WHERE id=$1 AND status='pending'`, id)
if err != nil {
return false, err
}
n, err := res.RowsAffected()
return n == 1, err
}
// RecordFindingRetestResult never accepts a finding ID: ownership comes from the
// runtime conversation. Identical retries are safe; a second verdict is refused.
func (d *DB) RecordFindingRetestResult(ctx context.Context, conversationID int64, verdict, summary, evidence string) error {
if verdict != "reproduced" && verdict != "fixed" && verdict != "inconclusive" {
return errors.New("verdict 必须为 reproduced / fixed / inconclusive")
}
summary, evidence = strings.TrimSpace(summary), strings.TrimSpace(evidence)
if summary == "" || evidence == "" {
return errors.New("summary 与 evidence 不能为空;无法确认时说明实际检查及阻塞原因")
}
if len(summary) > 16000 || len(evidence) > 128000 {
return errors.New("复测结论过长(summary ≤ 16KB,evidence ≤ 128KB)")
}
res, err := d.ExecContext(ctx, `UPDATE finding_retests SET verdict=$2,summary=$3,evidence=$4
WHERE conversation_id=$1 AND status='running' AND (verdict='' OR (verdict=$2 AND summary=$3 AND evidence=$4))`, conversationID, verdict, summary, evidence)
if err != nil {
return err
}
if n, err := res.RowsAffected(); err != nil {
return err
} else if n == 0 {
return ErrRetestNotRunning
}
return nil
}
// FinishFindingRetest seals the result. Cancellation/failure takes precedence
// over a staged verdict so an interrupted test cannot appear successfully fixed.
// Only a newly completed fixed verdict updates triage, in the same transaction.
func (d *DB) FinishFindingRetest(id int64, status, reason string) error {
if status != "completed" && status != "failed" && status != "stopped" {
return errors.New("invalid terminal retest status")
}
tx, err := d.Begin()
if err != nil {
return err
}
defer tx.Rollback()
// Lock the finding before the retest, matching creation and cascading deletion.
var findingID int64
err = tx.QueryRow(`SELECT f.id FROM findings f WHERE f.id=(SELECT finding_id FROM finding_retests WHERE id=$1) FOR UPDATE OF f`, id).Scan(&findingID)
if errors.Is(err, sql.ErrNoRows) {
return nil // Finding/retest already deleted.
}
if err != nil {
return err
}
var finalStatus, verdict string
err = tx.QueryRow(`UPDATE finding_retests SET
status=CASE WHEN $2='completed' AND verdict='' THEN 'failed' ELSE $2 END,
error=CASE WHEN $2='completed' AND verdict='' THEN '`+retestNoConclusionReason+`' ELSE $3 END,
finished_at=now() WHERE id=$1 AND status IN ('pending','running') RETURNING status,verdict`, id, status, reason).Scan(&finalStatus, &verdict)
if errors.Is(err, sql.ErrNoRows) {
return nil // A replay must not overwrite a later manual triage decision.
}
if err != nil {
return err
}
if finalStatus == "completed" && verdict == "fixed" {
// 走带通知的版本,与人工在详情页改状态共用同一套语义。
//
// 此前这里是裸的 UPDATE:复测判「已修复」时状态确实变了,但配了
// on_status_change 的渠道完全收不到推送——状态在界面上悄悄变了,
// 运维要打开平台才知道。状态更新与推送事件必须一起落库,
// SetFindingStatusTx 内部处理了「状态没变就不登记」等细节。
// 用 context.Background():本函数整条都是无 ctx 的旧风格(d.Begin()/
// tx.QueryRow/tx.Exec),没有可传递的取消信号,硬加一个 ctx 参数会
// 牵动 server 侧调用点与多处测试,超出本次改动的范围。
if _, _, _, _, err := SetFindingStatusTx(context.Background(), tx, findingID, FindingFixed); err != nil {
return err
}
}
return tx.Commit()
}
func (d *DB) RecoverFindingRetests() error {
_, err := d.Exec(`UPDATE finding_retests SET status='stopped', error='` + retestServiceRestartReason + `', finished_at=now() WHERE status IN ('pending','running')`)
return err
}
+37
View File
@@ -0,0 +1,37 @@
package db
import (
"testing"
"unicode"
)
// assertRetestReasonKorean 은 재검증 사유 상수가 한글을 포함하고 중국어 한자가 없음을
// 단언한다(F9). FinishFindingRetest·RecoverFindingRetests 의 SQL 리터럴 자리에 상수로
// 이어 붙이므로, 상수를 핀 고정하면 패널에 노출되는 사용자 문구도 함께 보호된다.
func assertRetestReasonKorean(t *testing.T, name, s string) {
t.Helper()
if s == "" {
t.Fatalf("%s: 빈 문자열", name)
}
hasHangul := false
for _, r := range s {
if unicode.Is(unicode.Han, r) {
t.Fatalf("%s: 중국어 한자가 남아 있습니다: %q", name, s)
}
if unicode.Is(unicode.Hangul, r) {
hasHangul = true
}
}
if !hasHangul {
t.Fatalf("%s: 한글이 없습니다: %q", name, s)
}
}
// TestFindingRetestReasonsLocalized 는 finding_retests.error 컬럼에 저장돼 재검증 패널
// (finding-retest-panel) 의 item.error 로 노출되는 종결 사유 두 상수가 한국어임을 단언한다.
// server/conversations.go 의 형제 사유(convRetest*)와 같은 컬럼·패널이라, 둘 중 하나만
// 한국어면 같은 패널에서 언어가 섞인다.
func TestFindingRetestReasonsLocalized(t *testing.T) {
assertRetestReasonKorean(t, "retestNoConclusionReason", retestNoConclusionReason)
assertRetestReasonKorean(t, "retestServiceRestartReason", retestServiceRestartReason)
}
+242
View File
@@ -0,0 +1,242 @@
package db
import (
"context"
"encoding/json"
"errors"
"strings"
"sync"
"testing"
)
func retestDB(t *testing.T) (*DB, int64) {
t.Helper()
d, err := Open(testDSN(t))
if err != nil {
t.Fatalf("open test database: %v", err)
}
fid, err := d.AddFinding(0, 0, "retest-test", "测试漏洞", "high", "original summary", "original evidence", "worker", nil)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
_, _ = d.Exec(`DELETE FROM conversations WHERE id IN (SELECT conversation_id FROM finding_retests WHERE finding_id=$1)`, fid)
_, _ = d.DeleteFinding(fid)
d.Close()
})
return d, fid
}
func TestRetestAtomicDeduplicationAndHistory(t *testing.T) {
d, fid := retestDB(t)
var wg sync.WaitGroup
var mu sync.Mutex
createdCount := 0
ids := make([]int64, 0, 8)
for range 8 {
wg.Go(func() {
r, c, created, err := d.CreateFindingRetest(t.Context(), fid, " 修复版本 v2 ")
if err != nil {
t.Error(err)
return
}
mu.Lock()
defer mu.Unlock()
ids = append(ids, r.ID)
if created {
createdCount++
if c == nil {
t.Error("created without conversation")
}
}
})
}
wg.Wait()
if createdCount != 1 || len(ids) != 8 {
t.Fatalf("created=%d ids=%v", createdCount, ids)
}
for _, id := range ids {
if id != ids[0] {
t.Fatal("duplicate active retests", ids)
}
}
rows, err := d.ListFindingRetests(fid)
if err != nil || len(rows) != 1 {
t.Fatalf("history=%v err=%v", rows, err)
}
r := rows[0]
if r.Snapshot != nil {
t.Fatal("history leaks large snapshot")
}
ctx := t.Context()
full, err := d.FindingRetestForConversation(ctx, *r.ConversationID)
if err != nil || !strings.Contains(string(full.Snapshot), "original evidence") {
t.Fatalf("snapshot=%+v err=%v", full, err)
}
var messages int
if err := d.QueryRow(`SELECT count(*) FROM conversation_activities WHERE conversation_id=$1 AND kind='user'`, *r.ConversationID).Scan(&messages); err != nil || messages != 1 {
t.Fatalf("messages=%d err=%v", messages, err)
}
_, _ = d.Exec(`UPDATE findings SET evidence='changed evidence' WHERE id=$1`, fid)
full, err = d.FindingRetestForConversation(ctx, *r.ConversationID)
if err != nil || strings.Contains(string(full.Snapshot), "changed evidence") {
t.Fatal("snapshot changed", err)
}
if err = d.RecordFindingRetestResult(ctx, *r.ConversationID, "fixed", "summary", "proof"); !errors.Is(err, ErrRetestNotRunning) {
t.Fatal("pending accepted result", err)
}
if ok, err := d.StartFindingRetest(ctx, r.ID); err != nil || !ok {
t.Fatalf("start=%t %v", ok, err)
}
if err := d.RecordFindingRetestResult(ctx, 0, "fixed", "summary", "proof"); !errors.Is(err, ErrRetestNotRunning) {
t.Fatal("unscoped write accepted", err)
}
for range 2 {
if err := d.RecordFindingRetestResult(ctx, *r.ConversationID, "fixed", "修复验证通过", "正常对照可用,原触发条件失效"); err != nil {
t.Fatal(err)
}
}
if err := d.RecordFindingRetestResult(ctx, *r.ConversationID, "reproduced", "different", "proof"); !errors.Is(err, ErrRetestNotRunning) {
t.Fatal("overwrote staged result", err)
}
if err := d.FinishFindingRetest(r.ID, "completed", ""); err != nil {
t.Fatal(err)
}
if err := d.RecordFindingRetestResult(ctx, *r.ConversationID, "fixed", "summary", "proof"); !errors.Is(err, ErrRetestNotRunning) {
t.Fatal("overwrote sealed result", err)
}
f, _ := d.GetFinding(fid)
if f.Status != FindingFixed || f.Evidence != "changed evidence" || f.Report != "" {
t.Fatal("fixed retest did not update only triage", f)
}
next, _, created, err := d.CreateFindingRetest(ctx, fid, "second")
if err != nil || !created || next.ID == r.ID {
t.Fatalf("new history=%+v created=%t err=%v", next, created, err)
}
if err := d.FinishFindingRetest(next.ID, "completed", ""); err != nil {
t.Fatal(err)
}
rows, _ = d.ListFindingRetests(fid)
if len(rows) != 2 || rows[0].Status != "failed" || rows[0].Error == "" || rows[1].Verdict != "fixed" {
t.Fatal("missing verdict treated as successful", rows)
}
}
func TestRetestDeletionAndRestart(t *testing.T) {
d, fid := retestDB(t)
r, c, _, err := d.CreateFindingRetest(t.Context(), fid, "delete")
if err != nil {
t.Fatal(err)
}
if err := d.DeleteConversation(c.ID); err != nil {
t.Fatal(err)
}
rows, _ := d.ListFindingRetests(fid)
if len(rows) != 1 || rows[0].ConversationID != nil || rows[0].Status != "stopped" {
t.Fatal(rows)
}
next, c2, created, err := d.CreateFindingRetest(t.Context(), fid, "restart")
if err != nil || !created {
t.Fatal("deleted session blocks retry", err)
}
if ok, err := d.StartFindingRetest(t.Context(), next.ID); err != nil || !ok {
t.Fatal(err)
}
if err := d.RecoverFindingRetests(); err != nil {
t.Fatal(err)
}
if err := d.FinishFindingRetest(next.ID, "completed", ""); err != nil {
t.Fatal(err)
}
full, _ := d.FindingRetestForConversation(t.Context(), c2.ID)
if full.Status != "stopped" || full.FinishedAt == nil {
t.Fatal("restart result overwritten", full)
}
if _, err := d.DeleteFinding(fid); err != nil {
t.Fatal(err)
}
var count int
_ = d.QueryRow(`SELECT count(*) FROM finding_retests WHERE id IN ($1,$2)`, r.ID, next.ID).Scan(&count)
if count != 0 {
t.Fatal("finding deletion did not cascade")
}
_ = d.DeleteConversation(c2.ID)
}
func TestRetestValidationAndRollback(t *testing.T) {
d, fid := retestDB(t)
ctx, cancel := context.WithCancel(t.Context())
cancel()
if _, _, _, err := d.CreateFindingRetest(ctx, fid, ""); err == nil {
t.Fatal("cancelled creation succeeded")
}
for _, args := range [][3]string{{"unknown", "summary", "proof"}, {"fixed", " ", "proof"}, {"fixed", "summary", ""}, {"fixed", strings.Repeat("x", 16001), "proof"}} {
if err := d.RecordFindingRetestResult(t.Context(), 0, args[0], args[1], args[2]); err == nil {
t.Fatal("bad result accepted", args[0])
}
}
r, c, _, err := d.CreateFindingRetest(t.Context(), fid, "snapshot")
if err != nil {
t.Fatal(err)
}
full, err := d.FindingRetestForConversation(t.Context(), c.ID)
var snap map[string]json.RawMessage
if err != nil || json.Unmarshal(full.Snapshot, &snap) != nil || len(snap["finding"]) == 0 {
t.Fatal("invalid snapshot", err)
}
if r.Notes != "snapshot" {
t.Fatal("notes lost")
}
}
func TestRetestFixedTriageOnlyAfterSuccessfulCompletion(t *testing.T) {
for _, tc := range []struct{ verdict, terminal, want string }{
{"fixed", "completed", FindingFixed},
{"fixed", "failed", FindingInProgress},
{"fixed", "stopped", FindingInProgress},
{"reproduced", "completed", FindingInProgress},
{"inconclusive", "completed", FindingInProgress},
{"", "completed", FindingInProgress},
} {
t.Run(tc.verdict+"/"+tc.terminal, func(t *testing.T) {
d, fid := retestDB(t)
if _, err := d.SetFindingStatus(fid, FindingInProgress); err != nil {
t.Fatal(err)
}
r, c, _, err := d.CreateFindingRetest(t.Context(), fid, "check triage")
if err != nil {
t.Fatal(err)
}
if _, err := d.StartFindingRetest(t.Context(), r.ID); err != nil {
t.Fatal(err)
}
if tc.verdict != "" {
if err := d.RecordFindingRetestResult(t.Context(), c.ID, tc.verdict, "summary", "proof"); err != nil {
t.Fatal(err)
}
}
f, err := d.GetFinding(fid)
if err != nil || f.Status != FindingInProgress {
t.Fatal("staged verdict changed triage", err, f)
}
if err := d.FinishFindingRetest(r.ID, tc.terminal, ""); err != nil {
t.Fatal(err)
}
f, err = d.GetFinding(fid)
if err != nil || f.Status != tc.want || f.Evidence != "original evidence" {
t.Fatalf("finding=%+v err=%v want=%s", f, err, tc.want)
}
// Re-delivering completion must not undo a later user decision.
if _, err := d.SetFindingStatus(fid, FindingIgnored); err != nil {
t.Fatal(err)
}
if err := d.FinishFindingRetest(r.ID, "completed", ""); err != nil {
t.Fatal(err)
}
f, err = d.GetFinding(fid)
if err != nil || f.Status != FindingIgnored {
t.Fatal("replayed completion overwrote triage", err, f)
}
})
}
}
+476
View File
@@ -0,0 +1,476 @@
package db
import (
"context"
"crypto/sha256"
"database/sql"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"strings"
"time"
"github.com/Autumn-27/artex/notify"
)
var (
ErrEvidenceConflict = errors.New("流量证据已变更,请刷新后重试")
ErrFindingNotFound = errors.New("漏洞不存在")
ErrEvidenceNotFound = errors.New("流量证据不存在")
)
// This lock covers the evidence filesystem as well as its SQL references. All
// processes sharing the database use it, including readers, exports and GC.
const findingEvidenceLockKey int64 = 7337741004
func (d *DB) WithEvidenceTx(ctx context.Context, fn func(*sql.Tx) error) error {
tx, err := d.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
if _, err = tx.ExecContext(ctx, `SELECT pg_advisory_xact_lock($1)`, findingEvidenceLockKey); err != nil {
return err
}
if err = fn(tx); err != nil {
return err
}
return tx.Commit()
}
type TrafficRef struct {
TrafficID string `json:"traffic_id"`
Role string `json:"role"`
Note string `json:"note"`
}
func NormalizeTrafficRefs(refs []TrafficRef) ([]TrafficRef, error) {
out := make([]TrafficRef, 0, len(refs))
seen := map[string]bool{}
for _, ref := range refs {
ref.TrafficID = strings.TrimSpace(ref.TrafficID)
if ref.TrafficID == "" {
return nil, errors.New("traffic_id 不能为空")
}
if ref.Role == "" {
ref.Role = "supporting"
}
if !ValidTrafficRole(ref.Role) {
return nil, fmt.Errorf("无效的流量用途 %q", ref.Role)
}
if !seen[ref.TrafficID] {
out = append(out, ref)
seen[ref.TrafficID] = true
}
}
return out, nil
}
func ValidTrafficRole(role string) bool {
return role == "baseline" || role == "proof" || role == "verification" || role == "supporting"
}
type TrafficEvidenceSnapshot struct {
ID string `json:"id"`
SourceTrafficID string `json:"source_traffic_id"`
CapturedAt int64 `json:"captured_at"`
URL string `json:"url"`
Method string `json:"method"`
Status int `json:"status"`
ContentType string `json:"content_type"`
ReqHead string `json:"req_head,omitempty"`
RespHead string `json:"resp_head,omitempty"`
ReqHash string `json:"req_hash"`
RespHash string `json:"resp_hash"`
ReqLen int64 `json:"req_len"`
RespLen int64 `json:"resp_len"`
CreatedAt time.Time `json:"created_at"`
}
type FindingTrafficBinding struct {
ID int64 `json:"id,string"`
FindingID int64 `json:"finding_id,string"`
SnapshotID string `json:"snapshot_id"`
Role string `json:"role"`
Note string `json:"note"`
Position int `json:"position"`
CreatedAt time.Time `json:"created_at"`
Snapshot TrafficEvidenceSnapshot `json:"snapshot"`
}
type FindingTraffic struct {
FindingID int64 `json:"finding_id,string"`
Version int64 `json:"version"`
ReportVersion int64 `json:"report_version"`
Bindings []FindingTrafficBinding `json:"bindings"`
}
type PreparedTrafficEvidence struct {
Ref TrafficRef
Snapshot TrafficEvidenceSnapshot
}
// Normalize makes the snapshot's text columns safe for PostgreSQL. URL and the
// head blocks come straight off the wire, so a target answering with a non-UTF-8
// header (a GBK `Content-Disposition: filename=…`, a NUL byte) would otherwise
// abort the INSERT and roll back the whole finding — losing a confirmed finding
// over a malformed response header. Applied before hashing so the ID always
// matches the bytes that actually land in the table.
func (s TrafficEvidenceSnapshot) Normalize() TrafficEvidenceSnapshot {
s.SourceTrafficID = utf8Clean(s.SourceTrafficID)
s.URL = utf8Clean(s.URL)
s.Method = utf8Clean(s.Method)
s.ContentType = utf8Clean(s.ContentType)
s.ReqHead = utf8Clean(s.ReqHead)
s.RespHead = utf8Clean(s.RespHead)
return s
}
func TrafficSnapshotID(snapshot TrafficEvidenceSnapshot) string {
snapshot = snapshot.Normalize()
snapshot.ID = ""
snapshot.CreatedAt = time.Time{}
raw, _ := json.Marshal(snapshot)
sum := sha256.Sum256(raw)
return hex.EncodeToString(sum[:])
}
// LockTaskEvidenceTx serializes writes with archive queueing (which locks the
// same task row). Once queued, its snapshot must not acquire new evidence.
func LockTaskEvidenceTx(tx *sql.Tx, taskID int64) error {
if taskID == 0 {
return nil
}
var deleted sql.NullTime
if err := tx.QueryRow(`SELECT deleted_at FROM tasks WHERE id=$1 FOR UPDATE`, taskID).Scan(&deleted); err != nil {
return err
}
if deleted.Valid {
return ErrTaskArchiveState
}
var state string
err := tx.QueryRow(`SELECT state FROM task_archives WHERE task_id=$1`, taskID).Scan(&state)
if err != nil && !errors.Is(err, sql.ErrNoRows) {
return err
}
if err == nil && state != ArchiveFailed {
return ErrTaskArchiveState
}
return nil
}
func LockFindingEvidenceTx(tx *sql.Tx, findingID int64, version *int64) error {
var taskID sql.NullInt64
if err := tx.QueryRow(`SELECT task_id FROM findings WHERE id=$1`, findingID).Scan(&taskID); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return ErrFindingNotFound
}
return err
}
if err := LockTaskEvidenceTx(tx, taskID.Int64); err != nil {
return err
}
var current int64
if err := tx.QueryRow(`SELECT evidence_version FROM findings WHERE id=$1 FOR UPDATE`, findingID).Scan(&current); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return ErrFindingNotFound
}
return err
}
if version != nil && *version != current {
return ErrEvidenceConflict
}
return nil
}
func InsertEvidenceSnapshotTx(tx *sql.Tx, s TrafficEvidenceSnapshot) error {
if s.ID != TrafficSnapshotID(s) {
return errors.New("证据快照元数据哈希不匹配")
}
// The ID was computed over the normalized form; store those same bytes.
id := s.ID
s = s.Normalize()
s.ID = id
if s.CreatedAt.IsZero() {
s.CreatedAt = time.Now().UTC()
}
_, err := tx.Exec(`INSERT INTO traffic_evidence_snapshots
(id,source_traffic_id,captured_at,url,method,status,content_type,req_head,resp_head,req_hash,resp_hash,req_len,resp_len,created_at)
VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14) ON CONFLICT(id) DO NOTHING`,
s.ID, s.SourceTrafficID, s.CapturedAt, s.URL, s.Method, s.Status, s.ContentType, s.ReqHead, s.RespHead, s.ReqHash, s.RespHash, s.ReqLen, s.RespLen, s.CreatedAt)
return err
}
func AddFindingTrafficTx(tx *sql.Tx, findingID int64, prepared []PreparedTrafficEvidence) error {
var position int
if err := tx.QueryRow(`SELECT COALESCE(MAX(position)+1,0) FROM finding_traffic_bindings WHERE finding_id=$1`, findingID).Scan(&position); err != nil {
return err
}
changed := false
for _, item := range prepared {
if err := InsertEvidenceSnapshotTx(tx, item.Snapshot); err != nil {
return err
}
if _, err := tx.Exec(`UPDATE traffic_evidence_snapshots SET unreferenced_at=NULL WHERE id=$1`, item.Snapshot.ID); err != nil {
return err
}
res, err := tx.Exec(`INSERT INTO finding_traffic_bindings(finding_id,snapshot_id,role,note,position)
VALUES($1,$2,$3,$4,$5) ON CONFLICT(finding_id,snapshot_id) DO NOTHING`, findingID, item.Snapshot.ID, item.Ref.Role, item.Ref.Note, position)
if err != nil {
return err
}
n, err := res.RowsAffected()
if err != nil {
return err
}
if n > 0 {
changed = true
position++
}
}
if changed {
return bumpEvidenceVersionTx(tx, findingID)
}
return nil
}
func bumpEvidenceVersionTx(tx *sql.Tx, findingID int64) error {
_, err := tx.Exec(`UPDATE findings SET evidence_version=evidence_version+1 WHERE id=$1`, findingID)
return err
}
func FindingTrafficTx(tx *sql.Tx, findingID int64) (*FindingTraffic, error) {
out := &FindingTraffic{FindingID: findingID, Bindings: []FindingTrafficBinding{}}
if err := tx.QueryRow(`SELECT evidence_version,report_evidence_version FROM findings WHERE id=$1`, findingID).Scan(&out.Version, &out.ReportVersion); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, ErrFindingNotFound
}
return nil, err
}
rows, err := tx.Query(`SELECT b.id,b.finding_id,b.snapshot_id,b.role,b.note,b.position,b.created_at,to_jsonb(s)
FROM finding_traffic_bindings b JOIN traffic_evidence_snapshots s ON s.id=b.snapshot_id WHERE b.finding_id=$1 ORDER BY b.position,b.id`, findingID)
if err != nil {
return nil, err
}
defer rows.Close()
for rows.Next() {
var b FindingTrafficBinding
var raw []byte
if err := rows.Scan(&b.ID, &b.FindingID, &b.SnapshotID, &b.Role, &b.Note, &b.Position, &b.CreatedAt, &raw); err != nil {
return nil, err
}
if err := json.Unmarshal(raw, &b.Snapshot); err != nil {
return nil, err
}
out.Bindings = append(out.Bindings, b)
}
return out, rows.Err()
}
func (d *DB) GetFindingTraffic(ctx context.Context, findingID int64) (out *FindingTraffic, err error) {
err = d.WithEvidenceTx(ctx, func(tx *sql.Tx) error { var e error; out, e = FindingTrafficTx(tx, findingID); return e })
return
}
func (d *DB) EditFindingTraffic(ctx context.Context, findingID, bindingID, version int64, role, note *string, remove bool, order []int64) error {
if role != nil && !ValidTrafficRole(*role) {
return errors.New("无效的流量用途")
}
return d.WithEvidenceTx(ctx, func(tx *sql.Tx) error {
if err := LockFindingEvidenceTx(tx, findingID, &version); err != nil {
return err
}
if order != nil {
current, err := FindingTrafficTx(tx, findingID)
if err != nil {
return err
}
if len(order) != len(current.Bindings) {
return ErrEvidenceConflict
}
ids := map[int64]bool{}
for _, b := range current.Bindings {
ids[b.ID] = true
}
for position, id := range order {
if !ids[id] {
return ErrEvidenceConflict
}
delete(ids, id)
if _, err := tx.Exec(`UPDATE finding_traffic_bindings SET position=$1 WHERE finding_id=$2 AND id=$3`, position, findingID, id); err != nil {
return err
}
}
} else {
var res sql.Result
var err error
if remove {
res, err = tx.Exec(`DELETE FROM finding_traffic_bindings WHERE finding_id=$1 AND id=$2`, findingID, bindingID)
} else {
res, err = tx.Exec(`UPDATE finding_traffic_bindings SET role=COALESCE($3,role),note=COALESCE($4,note) WHERE finding_id=$1 AND id=$2`, findingID, bindingID, role, note)
}
if err != nil {
return err
}
n, err := res.RowsAffected()
if err != nil {
return err
}
if n == 0 {
return ErrEvidenceNotFound
}
}
return bumpEvidenceVersionTx(tx, findingID)
})
}
type RecordFindingInput struct {
TaskID, ExplorationID, IntentID int64
VulnClass, Name, Severity, Summary, Evidence, Worker string
AssetIDs []int64
}
type RecordedFinding struct {
FindingID int64 `json:"finding_id,string"`
NodeID int64 `json:"finding_node_id,string"`
Traffic *FindingTraffic `json:"traffic"`
}
// ctx 由调用方传入本次事务所用的上下文(而非在内部取 context.Background):
// 事务内新加的推送事件写入同样应受调用方的取消与超时约束。
func RecordFindingTx(ctx context.Context, tx *sql.Tx, in RecordFindingInput, prepared []PreparedTrafficEvidence) (*RecordedFinding, error) {
if err := LockTaskEvidenceTx(tx, in.TaskID); err != nil {
return nil, err
}
if in.TaskID > 0 {
var expID int64
if err := tx.QueryRow(`SELECT exploration_id FROM tasks WHERE id=$1`, in.TaskID).Scan(&expID); err != nil {
return nil, err
}
if expID != in.ExplorationID {
return nil, errors.New("漏洞所属任务与探索记录不匹配")
}
}
if in.IntentID > 0 {
var ok bool
if err := tx.QueryRow(`SELECT EXISTS(SELECT 1 FROM exploration_nodes WHERE id=$1 AND exploration_id=$2 AND kind='intent')`, in.IntentID, in.ExplorationID).Scan(&ok); err != nil {
return nil, err
}
if !ok {
return nil, errors.New("intent_id 必须是本任务的意图(关联任务意图只读)")
}
}
payload, _ := json.Marshal(map[string]any{"vulnclass": in.VulnClass, "name": in.Name, "severity": in.Severity, "summary": in.Summary, "evidence": map[string]any{"by": in.Worker, "poc": in.Evidence}})
out := &RecordedFinding{}
if err := tx.QueryRow(`INSERT INTO exploration_nodes(exploration_id,kind,payload,priority,state,origin)
VALUES($1,'finding',$2,9,'confirmed',$3) RETURNING id`, in.ExplorationID, string(payload), in.Worker).Scan(&out.NodeID); err != nil {
return nil, err
}
for _, asset := range in.AssetIDs {
if _, err := tx.Exec(`INSERT INTO exploration_anchors(node_id,asset_id) VALUES($1,$2) ON CONFLICT DO NOTHING`, out.NodeID, asset); err != nil {
return nil, err
}
}
if in.IntentID > 0 {
if _, err := tx.Exec(`INSERT INTO exploration_edges(exploration_id,src_id,rel,dst_id) VALUES($1,$2,$3,$4)`, in.ExplorationID, in.IntentID, RelYields, out.NodeID); err != nil {
return nil, err
}
}
assets := in.AssetIDs
if assets == nil {
assets = []int64{}
}
raw, _ := json.Marshal(assets)
if err := tx.QueryRow(`INSERT INTO findings(task_id,node_id,vulnclass,name,severity,summary,evidence,worker,asset_ids)
VALUES(NULLIF($1,0),$2,$3,$4,$5,$6,$7,$8,$9) RETURNING id`, in.TaskID, out.NodeID, in.VulnClass, in.Name, in.Severity, in.Summary, in.Evidence, in.Worker, string(raw)).Scan(&out.FindingID); err != nil {
return nil, err
}
// 在**同一事务**里登记一条推送事件:提交即保证「漏洞落库」与「推送任务存在」
// 原子一致,不存在提交成功却没入队、消息永久丢失的窗口。
// 这里的失败被隔离在保存点上、不影响漏洞写入(见函数注释),因此忽略返回值。
RecordNotificationEventTx(ctx, tx, notify.EventFindingCreated, out.FindingID, notify.Snapshot{
Kind: notify.EventFindingCreated,
FindingID: out.FindingID,
TaskID: in.TaskID,
VulnClass: in.VulnClass,
Name: in.Name,
Severity: in.Severity,
Summary: in.Summary,
AssetIDs: assets,
})
if err := AddFindingTrafficTx(tx, out.FindingID, prepared); err != nil {
return nil, err
}
var err error
out.Traffic, err = FindingTrafficTx(tx, out.FindingID)
return out, err
}
// RecordFinding is the atomic legacy/no-recorder path used by tools and tests.
func (s *ExplorationStore) RecordFinding(ctx context.Context, in RecordFindingInput) (out *RecordedFinding, err error) {
in.ExplorationID = s.expID
err = s.db.WithEvidenceTx(ctx, func(tx *sql.Tx) error { var e error; out, e = RecordFindingTx(ctx, tx, in, nil); return e })
return
}
func (d *DB) FindingIDByNodeID(nodeID int64) (id int64, err error) {
err = d.QueryRow(`SELECT id FROM findings WHERE node_id=$1`, nodeID).Scan(&id)
return
}
// PopulateFindingTrafficIDs only enriches already-visible nodes. It performs no
// discovery or ID guessing, and leaves the legacy node ID unchanged.
func (s *ExplorationStore) PopulateFindingTrafficIDs(nodes []*Node) error {
byID := map[int64]*Node{}
var args []any
var placeholders []string
for _, n := range nodes {
if n == nil || n.Kind != KindFinding || byID[n.ID] != nil {
continue
}
byID[n.ID] = n
args = append(args, n.ID)
placeholders = append(placeholders, fmt.Sprintf("$%d", len(args)))
}
if len(args) == 0 {
return nil
}
rows, err := s.db.Query(`SELECT f.id,f.node_id,(SELECT count(*) FROM finding_traffic_bindings b WHERE b.finding_id=f.id) FROM findings f WHERE f.node_id IN (`+strings.Join(placeholders, ",")+`)`, args...)
if err != nil {
return err
}
defer rows.Close()
for rows.Next() {
var findingID, nodeID int64
var count int
if err := rows.Scan(&findingID, &nodeID, &count); err != nil {
return err
}
n := byID[nodeID]
n.FindingID, n.FindingNodeID, n.TrafficCount = findingID, nodeID, count
}
return rows.Err()
}
func (d *DB) SetFindingReportVersionByNodeID(ctx context.Context, nodeID int64, report string, version *int64) (n int64, err error) {
err = d.WithEvidenceTx(ctx, func(tx *sql.Tx) error {
var id int64
if err := tx.QueryRow(`SELECT id FROM findings WHERE node_id=$1`, nodeID).Scan(&id); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil
}
return err
}
if err := LockFindingEvidenceTx(tx, id, version); err != nil {
return err
}
res, err := tx.Exec(`UPDATE findings SET report=$2,report_evidence_version=COALESCE($3::bigint,CASE WHEN evidence_version=0 THEN 0 ELSE -1 END) WHERE id=$1`, id, report, version)
if err != nil {
return err
}
n, err = res.RowsAffected()
return err
})
return
}
+85
View File
@@ -0,0 +1,85 @@
package db
import (
"database/sql"
"encoding/json"
"errors"
"fmt"
)
func ArchiveEvidenceSnapshots(snapshot *TaskArchiveSnapshot) ([]TrafficEvidenceSnapshot, error) {
var out []TrafficEvidenceSnapshot
if rawRowCount(snapshot.Tables["traffic_evidence_snapshots"]) == 0 {
return out, nil
}
if snapshot.FormatVersion < 3 {
return nil, ErrTaskArchiveFormatMismatch
}
err := json.Unmarshal(snapshot.Tables["traffic_evidence_snapshots"], &out)
return out, err
}
func restoreFindingTrafficTx(tx *sql.Tx, snapshot *TaskArchiveSnapshot) error {
snapshots, err := ArchiveEvidenceSnapshots(snapshot)
if err != nil {
return err
}
allowed := map[string]bool{}
for _, v := range snapshots {
if allowed[v.ID] {
return errors.New("duplicate archived evidence snapshot")
}
allowed[v.ID] = true
if err = InsertEvidenceSnapshotTx(tx, v); err != nil {
return err
}
}
rows, err := decodeArchiveRows(snapshot.Tables["finding_traffic_bindings"])
if err != nil {
return err
}
if len(rows) > 0 && snapshot.FormatVersion < 3 {
return ErrTaskArchiveFormatMismatch
}
for _, row := range rows {
fid, ok := jsonInt64(row["finding_id"])
if !ok {
return errors.New("invalid archived evidence finding id")
}
sid, _ := row["snapshot_id"].(string)
role, _ := row["role"].(string)
if !allowed[sid] || !ValidTrafficRole(role) {
return errors.New("invalid archived evidence binding")
}
var owned bool
if err = tx.QueryRow(`SELECT EXISTS(SELECT 1 FROM findings WHERE id=$1 AND task_id=$2)`, fid, snapshot.TaskID).Scan(&owned); err != nil {
return err
}
if !owned {
return fmt.Errorf("evidence references finding outside archived task: %d", fid)
}
}
if len(rows) > 0 {
raw, _ := json.Marshal(rows)
if _, err = tx.Exec(`INSERT INTO finding_traffic_bindings SELECT * FROM json_populate_recordset(NULL::finding_traffic_bindings,$1::json)`, string(raw)); err != nil {
return err
}
_, err = tx.Exec(`UPDATE traffic_evidence_snapshots s SET unreferenced_at=NULL WHERE EXISTS(SELECT 1 FROM finding_traffic_bindings b WHERE b.snapshot_id=s.id)`)
}
return err
}
func normalizeArchivedFindingVersions(raw json.RawMessage) (json.RawMessage, error) {
rows, err := decodeArchiveRows(raw)
if err != nil {
return nil, err
}
for _, row := range rows {
for _, key := range []string{"evidence_version", "report_evidence_version"} {
if row[key] == nil {
row[key] = 0
}
}
}
return json.Marshal(rows)
}
+87
View File
@@ -0,0 +1,87 @@
package db
import (
"encoding/json"
"fmt"
"testing"
)
func TestFindingEvidenceLockNamespace(t *testing.T) {
for _, reserved := range []int64{7337741001, 7337741002, 7337741003} {
if findingEvidenceLockKey == reserved {
t.Fatal("evidence lock collides with migration/test/company lock")
}
}
}
func TestFindingTrafficLegacyArchiveDefaults(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Fatal(err)
}
defer d.Close()
if err = d.EnsureLLMRecordsTable(); err != nil {
t.Fatal(err)
}
if err = d.EnsureLLMUsageTable(); err != nil {
t.Fatal(err)
}
for _, version := range []int{1, 2} {
t.Run(fmt.Sprint(version), func(t *testing.T) {
task, err := d.CreateTask("legacy archive evidence defaults", "fixture", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
defer func() { d.Exec(`DELETE FROM task_archives WHERE task_id=$1`, task.ID); d.DeleteTask(task.ID) }()
f, err := d.Exploration(task.ExplorationID).RecordFinding(t.Context(), RecordFindingInput{TaskID: task.ID, ExplorationID: task.ExplorationID, Summary: "legacy", Severity: "low"})
if err != nil {
t.Fatal(err)
}
if err = d.SetPaused(task.ID, true); err != nil {
t.Fatal(err)
}
archive, err := d.QueueTaskArchive(task.ID)
if err != nil {
t.Fatal(err)
}
if _, err = d.ClaimTaskArchiveJob(t.Context()); err != nil {
t.Fatal(err)
}
snapshot, err := d.SnapshotTaskArchive(task.ID)
if err != nil {
t.Fatal(err)
}
if err = d.CompleteTaskArchive(archive.ID, snapshot, "/tmp/legacy-test.tar.zst", "fixture", 1, 1); err != nil {
t.Fatal(err)
}
snapshot.FormatVersion = version
delete(snapshot.Tables, "finding_traffic_bindings")
delete(snapshot.Tables, "traffic_evidence_snapshots")
var rows []map[string]any
if err = json.Unmarshal(snapshot.Tables["findings"], &rows); err != nil {
t.Fatal(err)
}
for _, row := range rows {
delete(row, "evidence_version")
delete(row, "report_evidence_version")
}
snapshot.Tables["findings"], _ = json.Marshal(rows)
if _, err = d.QueueTaskArchiveRestore(archive.ID); err != nil {
t.Fatal(err)
}
if _, err = d.ClaimTaskArchiveJob(t.Context()); err != nil {
t.Fatal(err)
}
if _, err = d.RestoreTaskArchive(archive.ID, snapshot, 0); err != nil {
t.Fatal(err)
}
got, err := d.GetFindingTraffic(t.Context(), f.FindingID)
if err != nil {
t.Fatal(err)
}
if got.Version != 0 || got.ReportVersion != 0 || len(got.Bindings) != 0 {
t.Fatal(got)
}
})
}
}
+794
View File
@@ -0,0 +1,794 @@
package db
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"sort"
"strconv"
"strings"
"time"
)
// DBFinding is a row in the standalone findings table. It persists across task
// deletion unless the caller explicitly requests related finding cleanup.
type DBFinding struct {
TrafficCount int
EvidenceVersion int64
ReportEvidenceVersion int64
TrafficBindings []FindingTrafficBinding // populated only for export
ID int64
TaskID *int64
NodeID *int64
VulnClass string
Name string // 漏洞名称(可读标题);为空时前端回退展示 VulnClass
Severity string
Summary string
Evidence string
Worker string
AssetIDs []int64
Status string
Report string // 详细报告(Markdown);仅 GetFinding 填充,列表查询不带
CreatedAt time.Time
TaskDescription string // populated via LEFT JOIN on tasks
}
// Finding triage states (findings.status).
const (
FindingPending = "pending" // 待处理
FindingInProgress = "in_progress" // 处理中
FindingConfirmed = "confirmed" // 已确认(真实漏洞,未修复)
FindingResolved = "resolved" // 已处理
FindingFixed = "fixed" // 已修复
FindingFalsePositive = "false_positive" // 误报
FindingIgnored = "ignored" // 忽略
FindingDuplicate = "duplicate" // 重复
FindingRiskAccepted = "risk_accepted" // 风险接受
)
// ValidFindingStatus reports whether s is a known triage state.
func ValidFindingStatus(s string) bool {
switch s {
case FindingPending, FindingInProgress, FindingConfirmed, FindingResolved, FindingFixed,
FindingFalsePositive, FindingIgnored, FindingDuplicate, FindingRiskAccepted:
return true
}
return false
}
// Finding severity levels (findings.severity).
const (
SeverityCritical = "critical" // 严重
SeverityHigh = "high" // 高
SeverityMedium = "medium" // 中
SeverityLow = "low" // 低
)
// ValidSeverity reports whether s is a known severity level.
func ValidSeverity(s string) bool {
switch s {
case SeverityCritical, SeverityHigh, SeverityMedium, SeverityLow:
return true
}
return false
}
// AddFinding inserts a finding into the standalone findings table. taskID and
// nodeID may be 0 (stored as NULL). name may be "" (frontend falls back to
// vulnclass). Returns the new finding id.
func (d *DB) AddFinding(taskID, nodeID int64, vulnclass, name, severity, summary, evidence, worker string, assetIDs []int64) (int64, error) {
aidsJSON, _ := json.Marshal(assetIDs)
if assetIDs == nil {
aidsJSON = []byte("[]")
}
var tid, nid *int64
if taskID > 0 {
tid = &taskID
}
if nodeID > 0 {
nid = &nodeID
}
var id int64
err := d.QueryRow(
`INSERT INTO findings (task_id, node_id, vulnclass, name, severity, summary, evidence, worker, asset_ids)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9) RETURNING id`,
tid, nid, vulnclass, name, severity, summary, evidence, worker, string(aidsJSON),
).Scan(&id)
return id, err
}
// findingSelectCols is the column list (with task_description join) every finding
// list query selects, so scanFinding stays in sync across callers.
const findingSelectCols = `f.id, f.task_id, f.node_id, f.vulnclass, COALESCE(f.name, ''), f.severity, f.summary,
f.evidence, f.worker, f.asset_ids, COALESCE(f.status, 'pending'), f.created_at,
COALESCE(t.description, '') AS task_description, f.evidence_version, f.report_evidence_version,
(SELECT count(*) FROM finding_traffic_bindings b WHERE b.finding_id=f.id)`
// scanFindings materializes rows selected via findingSelectCols.
func scanFindings(rows interface {
Next() bool
Scan(...any) error
Err() error
}) ([]*DBFinding, error) {
var out []*DBFinding
for rows.Next() {
f := &DBFinding{}
var aidsJSON string
if err := rows.Scan(&f.ID, &f.TaskID, &f.NodeID, &f.VulnClass, &f.Name, &f.Severity,
&f.Summary, &f.Evidence, &f.Worker, &aidsJSON, &f.Status, &f.CreatedAt, &f.TaskDescription, &f.EvidenceVersion, &f.ReportEvidenceVersion, &f.TrafficCount); err != nil {
return nil, err
}
_ = json.Unmarshal([]byte(aidsJSON), &f.AssetIDs)
out = append(out, f)
}
return out, rows.Err()
}
// ListFindings returns all findings (newest first), joined with task description.
// Kept for the dashboard's summary; the paginated 发现 page uses ListFindingsPage.
func (d *DB) ListFindings(limit int) ([]*DBFinding, error) {
if limit <= 0 {
limit = 500
}
rows, err := d.Query(`
SELECT `+findingSelectCols+`
FROM findings f
LEFT JOIN tasks t ON f.task_id = t.id
ORDER BY f.created_at DESC
LIMIT $1`, limit)
if err != nil {
return nil, err
}
defer rows.Close()
return scanFindings(rows)
}
// FindingFilter narrows a paginated findings query. Empty-string fields mean "no
// filter on that column". Sort is "severity" (severity desc, then newest) or
// anything else (newest first).
type FindingFilter struct {
Severity string // high | medium | low
Status string // pending | false_positive | ignored | resolved
VulnClass string
TaskID string // 任务 id(字符串形式;空/非法 = 不按任务筛选)
Query string // 名称/类型/摘要/证据/报告正文的模糊检索关键词
Sort string // "severity" | "time"
// AssetScope 是资产树的节点 key(a:<id> / c:<id> / r:<domain> / __none__),
// 选中一个节点等于选中它的整棵子树。空 = 不按资产筛选。
AssetScope string
// 下面三个由 applyAssetScope 从 AssetScope 解析而来,调用方不用设置。
assetIDs []int64 // 子树里所有资产 id
assetNone bool // 只要「未关联资产」的发现
assetMiss bool // 选中的节点在当前筛选下不存在 → 结果恒空
}
// FindingUnassignedTask is the task filter sentinel for findings whose task is
// absent. That includes rows created without a task and rows retained after their
// originating task was deleted (the findings FK is ON DELETE SET NULL).
const FindingUnassignedTask = "__unassigned__"
// where builds the WHERE clause (shared by the page and count queries) plus its
// positional args. All values are parameterized; Query also escapes ILIKE
// wildcards so user input is always matched literally.
func (f FindingFilter) where() (string, []any) {
var conds []string
var args []any
add := func(col, val string) {
if val == "" {
return
}
args = append(args, val)
conds = append(conds, fmt.Sprintf("f.%s = $%d", col, len(args)))
}
add("severity", f.Severity)
add("status", f.Status)
add("vulnclass", f.VulnClass)
// task_id 是 bigint 列,按整数比较(不能走上面的文本 add);空/非法值忽略。
if f.TaskID == FindingUnassignedTask {
conds = append(conds, "(f.task_id IS NULL OR t.id IS NULL)")
} else if tid, err := strconv.ParseInt(f.TaskID, 10, 64); err == nil && tid > 0 {
args = append(args, tid)
conds = append(conds, fmt.Sprintf("f.task_id = $%d", len(args)))
}
// 资产筛选:asset_ids 是 jsonb 数组,@> ANY(...) 能走 idx_findings_asset_ids。
switch {
case f.assetMiss:
conds = append(conds, "FALSE")
case f.assetNone:
// 「未关联资产」= asset_ids 为空,或者里面的 id 一个都不在 assets 表里
// (资产已被删除)。两类都进资产树的未关联桶,这里必须同样收下,否则桶上
// 的计数会大于点开后能查到的条数。
conds = append(conds, `(
jsonb_array_length(COALESCE(f.asset_ids, '[]'::jsonb)) = 0
OR NOT EXISTS (
SELECT 1 FROM jsonb_array_elements_text(f.asset_ids) e(v)
JOIN assets a ON a.id = e.v::bigint
)
)`)
case len(f.assetIDs) > 0:
args = append(args, assetIDContainments(f.assetIDs))
conds = append(conds, fmt.Sprintf("f.asset_ids @> ANY($%d::jsonb[])", len(args)))
}
if query := strings.TrimSpace(f.Query); query != "" {
escaped := strings.NewReplacer(`\`, `\\`, `%`, `\%`, `_`, `\_`).Replace(query)
args = append(args, "%"+escaped+"%")
placeholder := fmt.Sprintf("$%d", len(args))
conds = append(conds, fmt.Sprintf(`(
COALESCE(f.name, '') ILIKE %s ESCAPE '\' OR
f.vulnclass ILIKE %s ESCAPE '\' OR
f.summary ILIKE %s ESCAPE '\' OR
f.evidence ILIKE %s ESCAPE '\' OR
COALESCE(f.report, '') ILIKE %s ESCAPE '\'
)`, placeholder, placeholder, placeholder, placeholder, placeholder))
}
if len(conds) == 0 {
return "", args
}
return " WHERE " + strings.Join(conds, " AND "), args
}
// ListFindingsPage returns one page of findings matching the filter, plus the
// total count of matching rows (for the frontend pager). page is 1-based.
func (d *DB) ListFindingsPage(f FindingFilter, page, pageSize int) ([]*DBFinding, int, error) {
if page <= 0 {
page = 1
}
if pageSize <= 0 {
pageSize = 20
}
f, err := d.applyAssetScope(f)
if err != nil {
return nil, 0, err
}
where, args := f.where()
var total int
if err := d.QueryRow(`SELECT COUNT(*) FROM findings f LEFT JOIN tasks t ON f.task_id=t.id`+where, args...).Scan(&total); err != nil {
return nil, 0, err
}
// Avoid overflowing (page-1)*pageSize for an arbitrarily large page number.
// Once the requested page is beyond the exact count, no data query is needed.
if total == 0 || page > (total-1)/pageSize+1 {
return []*DBFinding{}, total, nil
}
order := "f.created_at DESC, f.id DESC"
if f.Sort == "severity" {
// critical > high > medium > low > 其它, then newest first.
order = `CASE f.severity WHEN 'critical' THEN 4 WHEN 'high' THEN 3 WHEN 'medium' THEN 2 WHEN 'low' THEN 1 ELSE 0 END DESC, f.created_at DESC, f.id DESC`
}
pageArgs := append(append([]any{}, args...), pageSize, (page-1)*pageSize)
q := fmt.Sprintf(`
SELECT %s
FROM findings f
LEFT JOIN tasks t ON f.task_id = t.id%s
ORDER BY %s
LIMIT $%d OFFSET $%d`, findingSelectCols, where, order, len(args)+1, len(args)+2)
rows, err := d.Query(q, pageArgs...)
if err != nil {
return nil, 0, err
}
defer rows.Close()
out, err := scanFindings(rows)
return out, total, err
}
// FindingGroup is one task-level bucket in the global findings view. TaskID is
// nil for both findings that never had a task and findings retained after task
// deletion; those records intentionally share one "unassigned/deleted" bucket.
type FindingGroup struct {
TaskID *int64 `json:"task_id"`
TaskName string `json:"task_name"` // 可选任务名称;空=未命名
TaskDescription string `json:"task_description"`
TaskStatus string `json:"task_status"`
Count int `json:"count"`
Critical int `json:"critical"`
High int `json:"high"`
Medium int `json:"medium"`
Low int `json:"low"`
LastFoundAt time.Time `json:"last_found_at"`
}
// ListFindingGroups returns a page of task groups matching the same filters as
// ListFindingsPage. The group count and finding count are independent totals so
// clients can page groups without losing the exact export/selection count.
func (d *DB) ListFindingGroups(f FindingFilter, page, pageSize int) ([]FindingGroup, int, int, error) {
if page <= 0 {
page = 1
}
if pageSize <= 0 {
pageSize = 10
}
f, err := d.applyAssetScope(f)
if err != nil {
return nil, 0, 0, err
}
where, args := f.where()
grouped := ` FROM findings f LEFT JOIN tasks t ON f.task_id=t.id` + where +
` GROUP BY t.id, t.name, t.description, t.status, t.paused, t.queued`
var groupTotal, findingTotal int
countQuery := `SELECT COUNT(*), COALESCE(SUM(finding_count),0) FROM (` +
`SELECT COUNT(*) AS finding_count` + grouped + `) grouped_findings`
if err := d.QueryRow(countQuery, args...).Scan(&groupTotal, &findingTotal); err != nil {
return nil, 0, 0, err
}
if groupTotal == 0 || page > (groupTotal-1)/pageSize+1 {
return []FindingGroup{}, groupTotal, findingTotal, nil
}
order := "MAX(f.created_at) DESC, t.id DESC NULLS LAST"
if f.Sort == "severity" {
order = `MAX(CASE f.severity WHEN 'critical' THEN 4 WHEN 'high' THEN 3 WHEN 'medium' THEN 2 WHEN 'low' THEN 1 ELSE 0 END) DESC, MAX(f.created_at) DESC, t.id DESC NULLS LAST`
}
pageArgs := append(append([]any{}, args...), pageSize, (page-1)*pageSize)
query := fmt.Sprintf(`SELECT t.id, COALESCE(t.name,''), COALESCE(t.description,''), COALESCE(
CASE
WHEN t.status IN ('done','failed','timeout') THEN t.status
WHEN t.queued THEN 'queued'
WHEN t.paused THEN 'paused'
ELSE t.status
END, ''),
COUNT(*),
COUNT(*) FILTER (WHERE f.severity='critical'),
COUNT(*) FILTER (WHERE f.severity='high'),
COUNT(*) FILTER (WHERE f.severity='medium'),
COUNT(*) FILTER (WHERE f.severity='low'),
MAX(f.created_at)%s
ORDER BY %s LIMIT $%d OFFSET $%d`, grouped, order, len(args)+1, len(args)+2)
rows, err := d.Query(query, pageArgs...)
if err != nil {
return nil, 0, 0, err
}
defer rows.Close()
groups := []FindingGroup{}
for rows.Next() {
var group FindingGroup
var taskID sql.NullInt64
if err := rows.Scan(&taskID, &group.TaskName, &group.TaskDescription, &group.TaskStatus, &group.Count,
&group.Critical, &group.High, &group.Medium, &group.Low, &group.LastFoundAt); err != nil {
return nil, 0, 0, err
}
if taskID.Valid {
id := taskID.Int64
group.TaskID = &id
}
groups = append(groups, group)
}
return groups, groupTotal, findingTotal, rows.Err()
}
// ErrFindingOriginUnavailable means a retained finding no longer has a live
// owning task and finding node from which a follow-up intent can be derived.
var ErrFindingOriginUnavailable = errors.New("finding origin is no longer available")
// AddFindingFollowUpIntent atomically creates a priority-10 human intent from a
// live finding node, copies that finding's asset anchors, records the
// finding --derived_from--> intent lineage edge, and persists its audit activity.
// The returned activity is the committed row and can be broadcast as-is without
// calling AppendActivity again.
func (s *ExplorationStore) AddFindingFollowUpIntent(findingID, findingNodeID int64, description string, audit Activity) (int64, Activity, error) {
tx, err := s.db.Begin()
if err != nil {
return 0, Activity{}, err
}
defer tx.Rollback()
var liveNodeID int64
err = tx.QueryRow(`SELECT n.id
FROM findings f
JOIN tasks t ON t.id=f.task_id
JOIN exploration_nodes n ON n.id=f.node_id AND n.exploration_id=t.exploration_id
WHERE f.id=$1 AND f.node_id=$2 AND t.exploration_id=$3 AND n.kind='finding'
FOR SHARE OF f, t, n`, findingID, findingNodeID, s.expID).Scan(&liveNodeID)
if err == sql.ErrNoRows {
return 0, Activity{}, ErrFindingOriginUnavailable
}
if err != nil {
return 0, Activity{}, err
}
anchors := []int64{}
anchorRows, err := tx.Query(`SELECT asset_id FROM exploration_anchors WHERE node_id=$1 ORDER BY asset_id`, liveNodeID)
if err != nil {
return 0, Activity{}, err
}
for anchorRows.Next() {
var assetID int64
if err := anchorRows.Scan(&assetID); err != nil {
anchorRows.Close()
return 0, Activity{}, err
}
anchors = append(anchors, assetID)
}
if err := anchorRows.Err(); err != nil {
anchorRows.Close()
return 0, Activity{}, err
}
if err := anchorRows.Close(); err != nil {
return 0, Activity{}, err
}
payload := map[string]any{
"summary": description,
"source_finding_id": findingID,
"source_finding_node_id": liveNodeID,
}
if len(anchors) > 0 {
payload["asset_ids"] = anchors
}
raw, err := json.Marshal(payload)
if err != nil {
return 0, Activity{}, err
}
var intentID int64
if err := tx.QueryRow(`INSERT INTO exploration_nodes(exploration_id,kind,payload,priority,state,origin)
VALUES ($1,'intent',$2,10,'open','human') RETURNING id`, s.expID, raw).Scan(&intentID); err != nil {
return 0, Activity{}, err
}
if _, err := tx.Exec(`INSERT INTO exploration_anchors(node_id,asset_id)
SELECT $1, asset_id FROM exploration_anchors WHERE node_id=$2
ON CONFLICT DO NOTHING`, intentID, liveNodeID); err != nil {
return 0, Activity{}, err
}
if _, err := tx.Exec(`INSERT INTO exploration_edges(exploration_id,src_id,rel,dst_id)
VALUES ($1,$2,$3,$4)`, s.expID, liveNodeID, RelDerivedFrom, intentID); err != nil {
return 0, Activity{}, err
}
audit.NodeID = &intentID
if summary := strings.TrimSpace(audit.Summary); summary != "" {
audit.Summary = fmt.Sprintf("%s #%d", summary, intentID)
} else {
audit.Summary = ""
}
metadata := audit.Metadata
if len(metadata) == 0 {
metadata = json.RawMessage(`{}`)
}
if err := tx.QueryRow(`
INSERT INTO activity(exploration_id, node_id, worker, kind, tool, tool_use_id, is_error, summary, detail, metadata, input_tokens, output_tokens, cache_read_tokens, cache_write_tokens)
VALUES ($1,$2,NULLIF($3,''),NULLIF($4,''),NULLIF($5,''),NULLIF($6,''),$7,NULLIF($8,''),NULLIF($9,''),$10,$11,$12,$13,$14)
RETURNING id, created_at`, s.expID, audit.NodeID, utf8Clean(audit.Worker), utf8Clean(audit.Kind), utf8Clean(audit.Tool), utf8Clean(audit.ToolUseID), audit.IsError,
utf8Clean(audit.Summary), utf8Clean(audit.Detail), metadata, audit.InputTokens, audit.OutputTokens, audit.CacheReadTokens, audit.CacheWriteTokens).
Scan(&audit.ID, &audit.CreatedAt); err != nil {
return 0, Activity{}, err
}
audit.Metadata = metadata
if err := tx.Commit(); err != nil {
return 0, Activity{}, err
}
return intentID, audit, nil
}
// ListFindingsForExport returns findings for the 发现 page 导出功能,携带完整
// report 字段、不分页。ids 非空时按这批 finding id 精确导出(勾选导出),忽略
// filter;ids 为空时按 filter 导出(导出当前筛选/全部)。结果按严重等级降序、
// 再按时间倒序,与「导出汇总报告」的分组顺序一致。
func (d *DB) ListFindingsForExport(f FindingFilter, ids []int64) ([]*DBFinding, error) {
const order = `ORDER BY CASE f.severity WHEN 'critical' THEN 4 WHEN 'high' THEN 3 WHEN 'medium' THEN 2 WHEN 'low' THEN 1 ELSE 0 END DESC, f.created_at DESC`
cols := findingSelectCols + `, COALESCE(f.report, '')`
var q string
var args []any
if len(ids) > 0 {
ph := make([]string, len(ids))
for i, id := range ids {
ph[i] = fmt.Sprintf("$%d", i+1)
args = append(args, id)
}
q = `SELECT ` + cols + `
FROM findings f
LEFT JOIN tasks t ON f.task_id = t.id
WHERE f.id IN (` + strings.Join(ph, ",") + `)
` + order
} else {
scoped, err := d.applyAssetScope(f)
if err != nil {
return nil, err
}
where, wargs := scoped.where()
q = `SELECT ` + cols + `
FROM findings f
LEFT JOIN tasks t ON f.task_id = t.id` + where + `
` + order
args = wargs
}
rows, err := d.Query(q, args...)
if err != nil {
return nil, err
}
defer rows.Close()
var out []*DBFinding
for rows.Next() {
f := &DBFinding{}
var aidsJSON string
if err := rows.Scan(&f.ID, &f.TaskID, &f.NodeID, &f.VulnClass, &f.Name, &f.Severity,
&f.Summary, &f.Evidence, &f.Worker, &aidsJSON, &f.Status, &f.CreatedAt,
&f.TaskDescription, &f.EvidenceVersion, &f.ReportEvidenceVersion, &f.TrafficCount, &f.Report); err != nil {
return nil, err
}
_ = json.Unmarshal([]byte(aidsJSON), &f.AssetIDs)
out = append(out, f)
}
return out, rows.Err()
}
// FindingStats is the whole-table aggregate powering the 发现 page's stat cards
// and vuln-class filter — computed server-side so it stays exact regardless of
// pagination.
type FindingStats struct {
Total int `json:"total"`
Pending int `json:"pending"`
Critical int `json:"critical"`
High int `json:"high"`
Medium int `json:"medium"`
Low int `json:"low"`
VulnClasses []string `json:"vulnclasses"`
Tasks []FindingTaskOption `json:"tasks"` // 有漏洞的任务(供「按任务」下拉)
}
// FindingTaskOption is one entry in the 发现 page's 任务 filter: a task that has at
// least one finding, with its description and finding count. Description is empty when
// the task has since been deleted (finding rows persist), so the frontend falls back to
// the id.
type FindingTaskOption struct {
ID int64 `json:"id"`
Name string `json:"name"` // 可选任务名称;空=未命名
Description string `json:"description"`
Count int `json:"count"`
}
// FindingStats returns whole-table counts (by severity + pending) and the sorted
// set of distinct vuln classes.
func (d *DB) FindingStats() (*FindingStats, error) {
st := &FindingStats{VulnClasses: []string{}, Tasks: []FindingTaskOption{}}
err := d.QueryRow(`SELECT
COUNT(*),
COUNT(*) FILTER (WHERE status = 'pending'),
COUNT(*) FILTER (WHERE severity = 'critical'),
COUNT(*) FILTER (WHERE severity = 'high'),
COUNT(*) FILTER (WHERE severity = 'medium'),
COUNT(*) FILTER (WHERE severity = 'low')
FROM findings`).Scan(&st.Total, &st.Pending, &st.Critical, &st.High, &st.Medium, &st.Low)
if err != nil {
return nil, err
}
rows, err := d.Query(`SELECT DISTINCT vulnclass FROM findings WHERE vulnclass <> '' ORDER BY vulnclass`)
if err != nil {
return nil, err
}
defer rows.Close()
for rows.Next() {
var vc string
if err := rows.Scan(&vc); err != nil {
return nil, err
}
st.VulnClasses = append(st.VulnClasses, vc)
}
if err := rows.Err(); err != nil {
return nil, err
}
// 任务下拉:有漏洞的任务,带描述(任务删除后为空,前端回退 id)和条数,最新有漏洞的排前。
trows, err := d.Query(`SELECT f.task_id, COALESCE(t.name, ''), COALESCE(t.description, ''), COUNT(*)
FROM findings f
LEFT JOIN tasks t ON f.task_id = t.id
WHERE f.task_id IS NOT NULL
GROUP BY f.task_id, t.name, t.description
ORDER BY MAX(f.created_at) DESC`)
if err != nil {
return nil, err
}
defer trows.Close()
for trows.Next() {
var opt FindingTaskOption
if err := trows.Scan(&opt.ID, &opt.Name, &opt.Description, &opt.Count); err != nil {
return nil, err
}
st.Tasks = append(st.Tasks, opt)
}
if err := trows.Err(); err != nil {
return nil, err
}
archived, err := d.archivedTaskAggregates()
if err != nil {
return nil, err
}
vulnclasses := make(map[string]bool, len(st.VulnClasses))
for _, vulnclass := range st.VulnClasses {
vulnclasses[vulnclass] = true
}
for _, aggregate := range archived {
cold := aggregate.FindingStats
st.Total += cold.Total
st.Pending += cold.Pending
st.Critical += cold.Critical
st.High += cold.High
st.Medium += cold.Medium
st.Low += cold.Low
for _, vulnclass := range cold.VulnClasses {
if vulnclass != "" {
vulnclasses[vulnclass] = true
}
}
}
st.VulnClasses = st.VulnClasses[:0]
for vulnclass := range vulnclasses {
st.VulnClasses = append(st.VulnClasses, vulnclass)
}
sort.Strings(st.VulnClasses)
return st, nil
}
// GetFinding returns a single finding row (with task_description joined and the
// full Markdown report), or nil when no row has that id. Unlike the list queries
// it also selects `report` — that column is only needed on the detail page.
func (d *DB) GetFinding(id int64) (*DBFinding, error) {
f := &DBFinding{}
var aidsJSON string
err := d.QueryRow(`SELECT `+findingSelectCols+`, COALESCE(f.report, '')
FROM findings f
LEFT JOIN tasks t ON f.task_id = t.id
WHERE f.id = $1`, id).Scan(
&f.ID, &f.TaskID, &f.NodeID, &f.VulnClass, &f.Name, &f.Severity,
&f.Summary, &f.Evidence, &f.Worker, &aidsJSON, &f.Status, &f.CreatedAt,
&f.TaskDescription, &f.EvidenceVersion, &f.ReportEvidenceVersion, &f.TrafficCount, &f.Report)
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
_ = json.Unmarshal([]byte(aidsJSON), &f.AssetIDs)
return f, nil
}
// DeleteFinding removes a finding entirely: the standalone findings row and its
// originating exploration node (kind='finding'), so it disappears from the findings
// list, the per-task 发现 Tab, and the exploration graph alike. Deleting the node
// cascades its edges + node_assets and nulls any activity referencing it. Returns
// rows affected (0 = no finding with that id).
func (d *DB) DeleteFinding(id int64) (n int64, err error) {
err = d.WithEvidenceTx(context.Background(), func(tx *sql.Tx) error {
if err := LockFindingEvidenceTx(tx, id, nil); err != nil {
if errors.Is(err, ErrFindingNotFound) {
return nil
}
return err
}
var nodeID sql.NullInt64
if err := tx.QueryRow(`DELETE FROM findings WHERE id=$1 RETURNING node_id`, id).Scan(&nodeID); err != nil {
return err
}
if nodeID.Valid {
if _, err := tx.Exec(`DELETE FROM exploration_nodes WHERE id=$1 AND kind='finding'`, nodeID.Int64); err != nil {
return err
}
}
n = 1
return nil
})
return
}
// DeleteFindingsByTask removes all findings rows of a task. The originating
// exploration finding nodes are cascade-deleted separately when the task's
// exploration subgraph is dropped. Returns rows deleted.
func (d *DB) DeleteFindingsByTask(taskID int64) (int64, error) {
res, err := d.Exec(`DELETE FROM findings WHERE task_id=$1`, taskID)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
// SetFindingStatus updates one finding's triage state. Returns rows affected.
//
// 底层 setter:只改状态、不登记推送事件。生产代码改状态请走
// SetFindingStatusWithNotify —— 直接调本函数会让「状态变更推送」静默失效。
// 保留它是为了让不关心通知的用例(参数校验、复测流程)能单独驱动状态。
func (d *DB) SetFindingStatus(id int64, status string) (int64, error) {
res, err := d.Exec(`UPDATE findings SET status=$1 WHERE id=$2`, status, id)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
// SetFindingReportByNodeID sets the Markdown report on the standalone finding row
// whose node_id matches — report_finding returns that node id, so an agent tool
// can address the finding it just created. Returns rows affected (0 when no row).
func (d *DB) SetFindingReportByNodeID(nodeID int64, report string) (int64, error) {
return d.SetFindingReportVersionByNodeID(context.Background(), nodeID, report, nil)
}
// setFindingCol updates one text column on the standalone finding row AND mirrors
// the new value into the originating exploration node's payload under jsonKey, so the
// per-task 发现 Tab (which reads the node payload, not this table) stays in sync.
// Returns rows affected (0 when no finding has that id); the node sync is best-effort.
// col and jsonKey MUST be trusted constants (they are interpolated into SQL) — never
// pass user input.
func (d *DB) setFindingCol(id int64, col, jsonKey, val string) (int64, error) {
var nodeID *int64
err := d.QueryRow(`UPDATE findings SET `+col+`=$1 WHERE id=$2 RETURNING node_id`, val, id).Scan(&nodeID)
if err == sql.ErrNoRows {
return 0, nil
}
if err != nil {
return 0, err
}
if nodeID != nil {
_, _ = d.Exec(`UPDATE exploration_nodes
SET payload = jsonb_set(payload, '{`+jsonKey+`}', to_jsonb($1::text))
WHERE id = $2`, val, *nodeID)
}
return 1, nil
}
// SetFindingSeverity updates one finding's severity (+ node payload sync). Returns
// rows affected (0 when no finding has that id).
func (d *DB) SetFindingSeverity(id int64, severity string) (int64, error) {
return d.setFindingCol(id, "severity", "severity", severity)
}
// SetFindingName updates one finding's 漏洞名称 (+ node payload sync). Empty name is
// allowed — the frontend falls back to the vuln class for display.
func (d *DB) SetFindingName(id int64, name string) (int64, error) {
return d.setFindingCol(id, "name", "name", name)
}
// SetFindingVulnClass updates one finding's 漏洞类别 (+ node payload sync).
func (d *DB) SetFindingVulnClass(id int64, vulnclass string) (int64, error) {
return d.setFindingCol(id, "vulnclass", "vulnclass", vulnclass)
}
// FindingMeta is the standalone-row data (id, triage state, anchored assets) the
// per-task view grafts onto its exploration-node findings.
type FindingMeta struct {
TrafficCount int
ID int64
Status string
AssetIDs []int64
}
// FindingMetaByNodeID maps a task's finding node ids to their standalone-row
// metadata (status + anchored asset ids) via the asset store, so callers holding
// only an AssetStore (e.g. the agent ToolSet) can reach it without a raw *DB.
func (a *AssetStore) FindingMetaByNodeID(taskID int64) (map[int64]FindingMeta, error) {
return a.db.FindingMetaByNodeID(taskID)
}
// FindingMetaByNodeID maps a task's finding node ids to their standalone-row
// metadata, so the per-task view (which reads exploration nodes) can show and
// edit the same status — and the same anchored assets — as the global 发现 page.
func (d *DB) FindingMetaByNodeID(taskID int64) (map[int64]FindingMeta, error) {
out := map[int64]FindingMeta{}
if taskID <= 0 {
return out, nil
}
rows, err := d.Query(`SELECT node_id, id, COALESCE(status,'pending'), asset_ids, (SELECT count(*) FROM finding_traffic_bindings b WHERE b.finding_id=findings.id) FROM findings
WHERE task_id=$1 AND node_id IS NOT NULL`, taskID)
if err != nil {
return out, err
}
defer rows.Close()
for rows.Next() {
var nid int64
var m FindingMeta
var aidsJSON string
if err := rows.Scan(&nid, &m.ID, &m.Status, &aidsJSON, &m.TrafficCount); err != nil {
return out, err
}
_ = json.Unmarshal([]byte(aidsJSON), &m.AssetIDs)
out[nid] = m
}
return out, rows.Err()
}
+519
View File
@@ -0,0 +1,519 @@
package db
import (
"encoding/json"
"errors"
"fmt"
"slices"
"strings"
"testing"
"time"
)
// TestDeleteFinding verifies删除漏洞 removes both the findings row and its
// originating exploration node (kind='finding').
func TestDeleteFinding(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
tk, err := d.CreateTask("删除漏洞测试", "目标", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
defer d.DeleteTask(tk.ID)
// seed a finding node in the task's exploration graph, then a findings row on it
es := d.Exploration(tk.ExplorationID)
nodeID, err := es.AddNode(KindFinding, map[string]any{"summary": "x", "severity": "high"}, 5, "confirmed", "worker", nil)
if err != nil {
t.Fatal(err)
}
fid, err := d.AddFinding(tk.ID, nodeID, "XSS", "反射型 XSS", "high", "summary", "poc", "worker", nil)
if err != nil {
t.Fatal(err)
}
n, err := d.DeleteFinding(fid)
if err != nil {
t.Fatal(err)
}
if n != 1 {
t.Fatalf("DeleteFinding rows: want 1, got %d", n)
}
if f, _ := d.GetFinding(fid); f != nil {
t.Fatalf("finding row should be gone, got %+v", f)
}
var cnt int
d.QueryRow(`SELECT count(*) FROM exploration_nodes WHERE id=$1`, nodeID).Scan(&cnt)
if cnt != 0 {
t.Fatalf("originating finding node should be deleted, still %d", cnt)
}
// deleting a non-existent finding is a no-op (0 rows), not an error
if n, err := d.DeleteFinding(fid); err != nil || n != 0 {
t.Fatalf("re-delete: want (0,nil), got (%d,%v)", n, err)
}
}
// TestFindingsPageAndStats exercises ListFindingsPage (filter/sort/paging) and
// FindingStats against the live dev PG. It tags its rows with a unique vulnclass
// so assertions are isolated from any pre-existing data, and cleans up after.
func TestFindingsPageAndStats(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
const vc = "__test_vc_pagination__"
// clean any leftovers from a prior aborted run, and clean up on exit
cleanup := func() { _, _ = d.Exec(`DELETE FROM findings WHERE vulnclass=$1`, vc) }
cleanup()
defer cleanup()
// Seed 6 findings under the marker vulnclass: 1 critical, 3 high, 2 low; 2 pending.
// The critical row carries a name to verify round-trip.
seed := []struct {
sev, status, name string
}{
{"critical", "pending", "严重漏洞标题"},
{"high", "resolved", ""},
{"high", "resolved", ""},
{"high", "pending", ""},
{"low", "resolved", ""},
{"low", "resolved", ""},
}
var ids []int64
for i, s := range seed {
id, err := d.AddFinding(0, 0, vc, s.name, s.sev, "summary", "poc", "tester", nil)
if err != nil {
t.Fatalf("AddFinding[%d]: %v", i, err)
}
if _, err := d.SetFindingStatus(id, s.status); err != nil {
t.Fatalf("SetFindingStatus[%d]: %v", i, err)
}
ids = append(ids, id)
}
// Filter by our vulnclass → exactly the 6 seeded rows, paged 2 per page.
p1, total, err := d.ListFindingsPage(FindingFilter{VulnClass: vc, Sort: "severity"}, 1, 2)
if err != nil {
t.Fatal(err)
}
if total != 6 {
t.Fatalf("total: want 6, got %d", total)
}
if len(p1) != 2 {
t.Fatalf("page1 size: want 2, got %d", len(p1))
}
// severity sort → critical first (with its name round-tripped), then high.
if p1[0].Severity != "critical" {
t.Fatalf("severity sort: want critical first, got %q", p1[0].Severity)
}
if p1[0].Name != "严重漏洞标题" {
t.Fatalf("name round-trip: want 严重漏洞标题, got %q", p1[0].Name)
}
if p1[1].Severity != "high" {
t.Fatalf("severity sort: want high second, got %q", p1[1].Severity)
}
// Combined filter: vulnclass + status=pending → 2 rows.
pend, total, err := d.ListFindingsPage(FindingFilter{VulnClass: vc, Status: FindingPending}, 1, 50)
if err != nil {
t.Fatal(err)
}
if total != 2 || len(pend) != 2 {
t.Fatalf("pending filter: want 2/2, got %d/%d", total, len(pend))
}
// Combined filter: vulnclass + severity=high → 3 rows.
_, total, err = d.ListFindingsPage(FindingFilter{VulnClass: vc, Severity: "high"}, 1, 50)
if err != nil {
t.Fatal(err)
}
if total != 3 {
t.Fatalf("high filter: want 3, got %d", total)
}
// Stats: whole-table, so assert our contribution is reflected (>=) and the
// marker vulnclass is present.
st, err := d.FindingStats()
if err != nil {
t.Fatal(err)
}
if st.Critical < 1 {
t.Fatalf("stats critical undercount: %+v", st)
}
if st.Total < 6 || st.High < 3 || st.Low < 2 || st.Pending < 2 {
t.Fatalf("stats undercount: %+v", st)
}
if !slices.Contains(st.VulnClasses, vc) {
t.Fatalf("stats vulnclasses missing %q", vc)
}
// GetFinding: single-row fetch round-trips id/name/severity.
one, err := d.GetFinding(ids[0])
if err != nil {
t.Fatal(err)
}
if one == nil || one.ID != ids[0] || one.Severity != "critical" || one.Name != "严重漏洞标题" {
t.Fatalf("GetFinding mismatch: %+v", one)
}
if one.Report != "" {
t.Fatalf("new finding report should be empty, got %q", one.Report)
}
// report column round-trips through GetFinding.
if _, err := d.Exec(`UPDATE findings SET report=$1 WHERE id=$2`, "# 报告\n正文", ids[0]); err != nil {
t.Fatal(err)
}
if one, _ = d.GetFinding(ids[0]); one.Report != "# 报告\n正文" {
t.Fatalf("report not read back: %q", one.Report)
}
for _, test := range []struct {
name string
query string
want int
}{
{name: "name", query: "严重漏洞标题", want: 1},
{name: "summary", query: "summary", want: 6},
{name: "evidence", query: "poc", want: 6},
{name: "report", query: "正文", want: 1},
{name: "case insensitive vulnclass", query: strings.ToUpper(vc), want: 6},
} {
t.Run("query_"+test.name, func(t *testing.T) {
matches, searchTotal, searchErr := d.ListFindingsPage(
FindingFilter{VulnClass: vc, Query: test.query}, 1, 20,
)
if searchErr != nil || searchTotal != test.want || len(matches) != test.want {
t.Fatalf("query %q: len=%d total=%d err=%v, want %d", test.query, len(matches), searchTotal, searchErr, test.want)
}
})
}
if _, err := d.Exec(`UPDATE findings SET name=$1 WHERE id=$2`, "literal %_ marker", ids[1]); err != nil {
t.Fatal(err)
}
matches, searchTotal, err := d.ListFindingsPage(FindingFilter{VulnClass: vc, Query: "%_"}, 1, 20)
if err != nil || searchTotal != 1 || len(matches) != 1 || matches[0].ID != ids[1] {
t.Fatalf("query wildcards must be literal: matches=%+v total=%d err=%v", matches, searchTotal, err)
}
if miss, err := d.GetFinding(-1); err != nil || miss != nil {
t.Fatalf("GetFinding(-1): want nil,nil got %+v,%v", miss, err)
}
// SetFindingSeverity: standalone row updates; 0 rows for unknown id.
if n, err := d.SetFindingSeverity(ids[0], "high"); err != nil || n != 1 {
t.Fatalf("SetFindingSeverity: want 1,nil got %d,%v", n, err)
}
one, _ = d.GetFinding(ids[0])
if one.Severity != "high" {
t.Fatalf("severity not updated: %q", one.Severity)
}
if n, err := d.SetFindingSeverity(-1, "low"); err != nil || n != 0 {
t.Fatalf("SetFindingSeverity(-1): want 0,nil got %d,%v", n, err)
}
}
func TestListFindingsPageUsesStableIDTieBreaker(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
vc := fmt.Sprintf("__test_finding_stable_page_%d__", time.Now().UnixNano())
defer d.Exec(`DELETE FROM findings WHERE vulnclass=$1`, vc)
ids := make([]int64, 0, 7)
for i := 0; i < 7; i++ {
id, addErr := d.AddFinding(0, 0, vc, "", SeverityHigh, fmt.Sprintf("finding %d", i), "", "test", nil)
if addErr != nil {
t.Fatalf("AddFinding[%d]: %v", i, addErr)
}
ids = append(ids, id)
}
sharedCreatedAt := time.Date(2026, time.August, 21, 8, 30, 0, 0, time.UTC)
if _, err := d.Exec(`UPDATE findings SET created_at=$1 WHERE vulnclass=$2`, sharedCreatedAt, vc); err != nil {
t.Fatal(err)
}
want := make([]int64, len(ids))
for i := range ids {
want[i] = ids[len(ids)-1-i]
}
for _, sort := range []string{"time", "severity"} {
t.Run(sort, func(t *testing.T) {
var got []int64
for page := 1; page <= 3; page++ {
items, total, pageErr := d.ListFindingsPage(FindingFilter{VulnClass: vc, Sort: sort}, page, 3)
if pageErr != nil {
t.Fatal(pageErr)
}
if total != len(ids) {
t.Fatalf("page %d total=%d, want %d", page, total, len(ids))
}
for _, item := range items {
got = append(got, item.ID)
}
}
if !slices.Equal(got, want) {
t.Fatalf("same-timestamp pagination was unstable: got=%v want=%v", got, want)
}
})
}
}
func TestFindingGroupsAndUnassignedPaging(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
vc := fmt.Sprintf("__test_finding_groups_%d__", time.Now().UnixNano())
cleanupFindings := func() { _, _ = d.Exec(`DELETE FROM findings WHERE vulnclass=$1`, vc) }
defer cleanupFindings()
taskA, err := d.CreateTask("group A", "goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
defer d.DeleteTask(taskA.ID)
taskB, err := d.CreateTask("group B", "goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
defer d.DeleteTask(taskB.ID)
seed := []struct {
taskID int64
severity string
status string
}{
{taskA.ID, SeverityCritical, FindingPending},
{taskA.ID, SeverityHigh, FindingResolved},
{taskB.ID, SeverityLow, FindingPending},
{0, SeverityMedium, FindingPending},
}
for i, item := range seed {
id, addErr := d.AddFinding(item.taskID, 0, vc, "", item.severity, fmt.Sprintf("summary %d", i), "", "test", nil)
if addErr != nil {
t.Fatalf("AddFinding[%d]: %v", i, addErr)
}
if _, setErr := d.SetFindingStatus(id, item.status); setErr != nil {
t.Fatalf("SetFindingStatus[%d]: %v", i, setErr)
}
}
matchedGroups, matchedGroupTotal, matchedFindingTotal, err := d.ListFindingGroups(
FindingFilter{VulnClass: vc, Query: "summary 0"}, 1, 10,
)
if err != nil || matchedGroupTotal != 1 || matchedFindingTotal != 1 || len(matchedGroups) != 1 ||
matchedGroups[0].TaskID == nil || *matchedGroups[0].TaskID != taskA.ID {
t.Fatalf("query-filtered groups: %+v groups=%d findings=%d err=%v",
matchedGroups, matchedGroupTotal, matchedFindingTotal, err)
}
// Retained findings from deleted tasks join the same bucket as findings that
// were created without any task.
if err := d.DeleteTask(taskB.ID); err != nil {
t.Fatal(err)
}
page, groupTotal, findingTotal, err := d.ListFindingGroups(FindingFilter{VulnClass: vc, Sort: "severity"}, 1, 1)
if err != nil {
t.Fatal(err)
}
if len(page) != 1 || groupTotal != 2 || findingTotal != 4 {
t.Fatalf("page totals: items=%d groups=%d findings=%d", len(page), groupTotal, findingTotal)
}
groups, _, _, err := d.ListFindingGroups(FindingFilter{VulnClass: vc}, 1, 10)
if err != nil {
t.Fatal(err)
}
var live, unassigned *FindingGroup
for i := range groups {
if groups[i].TaskID == nil {
unassigned = &groups[i]
} else if *groups[i].TaskID == taskA.ID {
live = &groups[i]
}
}
if live == nil || live.Count != 2 || live.Critical != 1 || live.High != 1 || live.TaskDescription != "group A" {
t.Fatalf("live group mismatch: %+v", live)
}
if unassigned == nil || unassigned.Count != 2 || unassigned.Medium != 1 || unassigned.Low != 1 {
t.Fatalf("unassigned group mismatch: %+v", unassigned)
}
groupForTask := func(items []FindingGroup, taskID int64) *FindingGroup {
for i := range items {
if items[i].TaskID != nil && *items[i].TaskID == taskID {
return &items[i]
}
}
return nil
}
if err := d.SetPaused(taskA.ID, true); err != nil {
t.Fatal(err)
}
groups, _, _, err = d.ListFindingGroups(FindingFilter{VulnClass: vc}, 1, 10)
if err != nil {
t.Fatal(err)
}
if group := groupForTask(groups, taskA.ID); group == nil || group.TaskStatus != "paused" {
t.Fatalf("paused task group status mismatch: %+v", groups)
}
if err := d.SetPaused(taskA.ID, false); err != nil {
t.Fatal(err)
}
if err := d.Enqueue(taskA.ID, "resume"); err != nil {
t.Fatal(err)
}
groups, _, _, err = d.ListFindingGroups(FindingFilter{VulnClass: vc}, 1, 10)
if err != nil {
t.Fatal(err)
}
if group := groupForTask(groups, taskA.ID); group == nil || group.TaskStatus != "queued" {
t.Fatalf("queued task group status mismatch: %+v", groups)
}
if err := d.SetStatus(taskA.ID, "done"); err != nil {
t.Fatal(err)
}
groups, _, _, err = d.ListFindingGroups(FindingFilter{VulnClass: vc}, 1, 10)
if err != nil {
t.Fatal(err)
}
if group := groupForTask(groups, taskA.ID); group == nil || group.TaskStatus != "done" {
t.Fatalf("terminal task status must win over queued flag: %+v", groups)
}
orphans, total, err := d.ListFindingsPage(FindingFilter{VulnClass: vc, TaskID: FindingUnassignedTask}, 1, 10)
if err != nil || total != 2 || len(orphans) != 2 {
t.Fatalf("unassigned page: len=%d total=%d err=%v", len(orphans), total, err)
}
maxPage := int(^uint(0) >> 1)
farFindings, farTotal, err := d.ListFindingsPage(FindingFilter{VulnClass: vc}, maxPage, 200)
if err != nil || farTotal != len(seed) || len(farFindings) != 0 {
t.Fatalf("far finding page: len=%d total=%d err=%v", len(farFindings), farTotal, err)
}
farGroups, farGroupTotal, farFindingTotal, err := d.ListFindingGroups(FindingFilter{VulnClass: vc}, maxPage, 100)
if err != nil || farGroupTotal != 2 || farFindingTotal != len(seed) || len(farGroups) != 0 {
t.Fatalf("far group page: len=%d groups=%d findings=%d err=%v",
len(farGroups), farGroupTotal, farFindingTotal, err)
}
filtered, filteredGroups, filteredFindings, err := d.ListFindingGroups(
FindingFilter{VulnClass: vc, Status: FindingResolved}, 1, 10,
)
if err != nil || filteredGroups != 1 || filteredFindings != 1 || len(filtered) != 1 || filtered[0].High != 1 {
t.Fatalf("filtered groups: %+v groups=%d findings=%d err=%v", filtered, filteredGroups, filteredFindings, err)
}
}
func TestAddFindingFollowUpIntent(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
task, err := d.CreateTask("finding follow-up", "goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
defer d.DeleteTask(task.ID)
assetID, err := d.Assets().UpsertRootDomain(UpsertRootDomainReq{
Domain: fmt.Sprintf("follow-up-%d.example.test", time.Now().UnixNano()),
TaskID: task.ID,
})
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM assets WHERE id=$1`, assetID)
store := d.Exploration(task.ExplorationID)
findingNodeID, err := store.AddNode(KindFinding, map[string]any{"summary": "source"}, 5, "confirmed", "worker", []int64{assetID})
if err != nil {
t.Fatal(err)
}
findingID, err := d.AddFinding(task.ID, findingNodeID, "test", "source", SeverityHigh, "source", "", "worker", []int64{assetID})
if err != nil {
t.Fatal(err)
}
auditInput := Activity{Worker: "system", Kind: "text", Summary: "사용자가 제출한 취약점 심화 익스플로잇 의도", Detail: "验证可利用性并形成证据链"}
intentID, audit, err := store.AddFindingFollowUpIntent(findingID, findingNodeID, "验证可利用性并形成证据链", auditInput)
if err != nil {
t.Fatal(err)
}
secondID, _, err := store.AddFindingFollowUpIntent(findingID, findingNodeID, "从另一条路径深入", Activity{
Worker: "system", Kind: "text", Summary: "사용자가 제출한 취약점 심화 익스플로잇 의도", Detail: "从另一条路径深入",
})
if err != nil {
t.Fatal(err)
}
if secondID == intentID {
t.Fatal("repeated follow-up submissions must create distinct intents")
}
expectedAuditSummary := fmt.Sprintf("%s #%d", auditInput.Summary, intentID)
if audit.ID <= 0 || audit.NodeID == nil || *audit.NodeID != intentID || audit.CreatedAt.IsZero() || audit.Summary != expectedAuditSummary {
t.Fatalf("persisted audit mismatch: %+v", audit)
}
node, err := store.GetNode(intentID)
if err != nil || node == nil {
t.Fatalf("GetNode: node=%+v err=%v", node, err)
}
if node.Kind != KindIntent || node.Priority != 10 || node.State != "open" || node.Origin != "human" {
t.Fatalf("follow-up intent metadata: %+v", node)
}
var anchorCount, edgeCount int
if err := d.QueryRow(`SELECT COUNT(*) FROM exploration_anchors WHERE node_id=$1 AND asset_id=$2`, intentID, assetID).Scan(&anchorCount); err != nil {
t.Fatal(err)
}
if err := d.QueryRow(`SELECT COUNT(*) FROM exploration_edges
WHERE exploration_id=$1 AND src_id=$2 AND rel=$3 AND dst_id=$4`,
task.ExplorationID, findingNodeID, RelDerivedFrom, intentID).Scan(&edgeCount); err != nil {
t.Fatal(err)
}
if anchorCount != 1 || edgeCount != 1 {
t.Fatalf("follow-up lineage: anchors=%d edges=%d", anchorCount, edgeCount)
}
var persistedAuditCount int
if err := d.QueryRow(`SELECT COUNT(*) FROM activity
WHERE id=$1 AND exploration_id=$2 AND node_id=$3 AND worker='system' AND summary=$4`,
audit.ID, task.ExplorationID, intentID, expectedAuditSummary).Scan(&persistedAuditCount); err != nil {
t.Fatal(err)
}
if persistedAuditCount != 1 {
t.Fatalf("atomic audit count=%d, want 1", persistedAuditCount)
}
var nodesBefore, activitiesBefore int
if err := d.QueryRow(`SELECT COUNT(*) FROM exploration_nodes WHERE exploration_id=$1`, task.ExplorationID).Scan(&nodesBefore); err != nil {
t.Fatal(err)
}
if err := d.QueryRow(`SELECT COUNT(*) FROM activity WHERE exploration_id=$1`, task.ExplorationID).Scan(&activitiesBefore); err != nil {
t.Fatal(err)
}
if _, _, err := store.AddFindingFollowUpIntent(findingID, findingNodeID, "must roll back", Activity{
Worker: "system", Kind: "text", Summary: "invalid audit", Metadata: json.RawMessage(`{`),
}); err == nil {
t.Fatal("invalid activity metadata should fail the transaction")
}
var nodesAfter, activitiesAfter int
if err := d.QueryRow(`SELECT COUNT(*) FROM exploration_nodes WHERE exploration_id=$1`, task.ExplorationID).Scan(&nodesAfter); err != nil {
t.Fatal(err)
}
if err := d.QueryRow(`SELECT COUNT(*) FROM activity WHERE exploration_id=$1`, task.ExplorationID).Scan(&activitiesAfter); err != nil {
t.Fatal(err)
}
if nodesAfter != nodesBefore || activitiesAfter != activitiesBefore {
t.Fatalf("activity failure did not roll back graph: nodes %d->%d activities %d->%d",
nodesBefore, nodesAfter, activitiesBefore, activitiesAfter)
}
if _, err := d.Exec(`DELETE FROM exploration_nodes WHERE id=$1`, findingNodeID); err != nil {
t.Fatal(err)
}
if _, _, err := store.AddFindingFollowUpIntent(findingID, findingNodeID, "should fail", auditInput); !errors.Is(err, ErrFindingOriginUnavailable) {
t.Fatalf("missing source node: want ErrFindingOriginUnavailable, got %v", err)
}
}
+32
View File
@@ -0,0 +1,32 @@
package db
import "fmt"
// DiscardOpenIntent compensates a follow-up creation when task admission fails.
// The task execution gate must still be held by the caller, so the intent cannot
// be claimed between this check and deletion. Edges and anchors cascade with the
// node; activity is deleted explicitly because its node FK otherwise becomes NULL.
func (s *ExplorationStore) DiscardOpenIntent(id int64) error {
tx, err := s.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
if _, err := tx.Exec(`DELETE FROM activity WHERE exploration_id=$1 AND node_id=$2`, s.expID, id); err != nil {
return err
}
res, err := tx.Exec(`DELETE FROM exploration_nodes
WHERE id=$1 AND exploration_id=$2 AND kind='intent' AND state='open'`, id, s.expID)
if err != nil {
return err
}
removed, err := res.RowsAffected()
if err != nil {
return err
}
if removed != 1 {
return fmt.Errorf("open intent %d was not available for admission rollback", id)
}
return tx.Commit()
}
+285
View File
@@ -0,0 +1,285 @@
package db
import (
"fmt"
"sync"
"sync/atomic"
"testing"
"time"
)
func TestCompareAndSetIntentStateAllowsSingleWinner(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
expID, err := d.CreateExploration("intent CAS", "only one controller wins")
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM explorations WHERE id=$1`, expID)
store := d.Exploration(expID)
intentID, err := store.AddIntent(map[string]any{"summary": "controlled"}, 1, nil, "planner")
if err != nil {
t.Fatal(err)
}
if claimed, err := store.ClaimIntent(intentID, "worker"); err != nil || !claimed {
t.Fatalf("claim: claimed=%v err=%v", claimed, err)
}
const contenders = 12
var winners atomic.Int32
var wg sync.WaitGroup
start := make(chan struct{})
for range contenders {
wg.Add(1)
go func() {
defer wg.Done()
<-start
changed, transitionErr := store.CompareAndSetIntentState(intentID, "running", "paused")
if transitionErr != nil {
t.Errorf("transition: %v", transitionErr)
return
}
if changed {
winners.Add(1)
}
}()
}
close(start)
wg.Wait()
if got := winners.Load(); got != 1 {
t.Fatalf("CAS winners=%d, want 1", got)
}
node, err := store.GetNode(intentID)
if err != nil || node == nil || node.State != "paused" {
t.Fatalf("node=%+v err=%v, want paused", node, err)
}
}
func TestCancelIntentPreservesTokenMeteringWithoutDoubleCount(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
expID, err := d.CreateExploration("cancel token rollup", "preserve consumed tokens")
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM explorations WHERE id=$1`, expID)
store := d.Exploration(expID)
intentID, err := store.AddIntent(map[string]any{"summary": "cancelled"}, 1, nil, "planner")
if err != nil {
t.Fatal(err)
}
if claimed, err := store.ClaimIntent(intentID, "work#1"); err != nil || !claimed {
t.Fatalf("claim: claimed=%v err=%v", claimed, err)
}
baselineDaily, err := d.TokenDailyAll(30)
if err != nil {
t.Fatal(err)
}
baseline := dailyTokenBuckets(baselineDaily)
appendUsage := func(kind string, at time.Time, input, output, read, write *int) {
t.Helper()
activityID, appendErr := store.AppendActivity(Activity{
NodeID: &intentID, Worker: "work#1", Kind: kind,
InputTokens: input, OutputTokens: output,
CacheReadTokens: read, CacheWriteTokens: write,
})
if appendErr != nil {
t.Fatal(appendErr)
}
if _, updateErr := d.Exec(`UPDATE activity SET created_at=$1 WHERE id=$2`, at, activityID); updateErr != nil {
t.Fatal(updateErr)
}
}
values := func(input, output, read, write int) (*int, *int, *int, *int) {
return &input, &output, &read, &write
}
now := time.Now().UTC()
atUTCNoon := func(daysAgo int) time.Time {
at := now.AddDate(0, 0, -daysAgo)
return time.Date(at.Year(), at.Month(), at.Day(), 12, 0, 0, 0, time.UTC)
}
authoritativeDay := atUTCNoon(6)
fallbackDay := atUTCNoon(4)
trailingDay := atUTCNoon(2)
// Run 1 has an authoritative result; its preceding cumulative frame must not
// be counted a second time or move usage to the frame's earlier date.
i, o, r, w := values(100, 100, 100, 100)
appendUsage("usage", atUTCNoon(7), i, o, r, w)
i, o, r, w = values(10, 11, 12, 13)
appendUsage("result", authoritativeDay, i, o, r, w)
// Run 2 models a legacy failed result that omitted usage. Its fallback belongs
// to the result's date, not the cumulative frame's date.
i, o, r, w = values(20, 21, 22, 23)
appendUsage("usage", atUTCNoon(5), i, o, r, w)
appendUsage("result", fallbackDay, nil, nil, nil, nil)
// Run 3 was interrupted before a terminal result was persisted.
i, o, r, w = values(30, 31, 32, 33)
appendUsage("usage", trailingDay, i, o, r, w)
beforeCancelDaily, err := d.TokenDailyAll(30)
if err != nil {
t.Fatal(err)
}
beforeCancel := dailyTokenBuckets(beforeCancelDaily)
assertDailyTokenDelta(t, beforeCancel, baseline, authoritativeDay, 10, 11, 12)
assertDailyTokenDelta(t, beforeCancel, baseline, fallbackDay, 0, 0, 0)
assertDailyTokenDelta(t, beforeCancel, baseline, trailingDay, 0, 0, 0)
if err := store.SetIntentState(intentID, "paused"); err != nil {
t.Fatal(err)
}
cleanup, err := store.CancelIntent(intentID)
if err != nil {
t.Fatal(err)
}
if cleanup.Activities != 5 {
t.Fatalf("deleted activities=%d, want 5", cleanup.Activities)
}
total, err := store.TokenTotal()
if err != nil {
t.Fatal(err)
}
assertTokenUsage(t, total, 60, 63, 66, 69)
afterCancelDaily, err := d.TokenDailyAll(30)
if err != nil {
t.Fatal(err)
}
afterCancel := dailyTokenBuckets(afterCancelDaily)
assertDailyTokenDelta(t, afterCancel, baseline, authoritativeDay, 10, 11, 12)
assertDailyTokenDelta(t, afterCancel, baseline, fallbackDay, 20, 21, 22)
assertDailyTokenDelta(t, afterCancel, baseline, trailingDay, 30, 31, 32)
assertDailyTokenDelta(t, afterCancel, baseline, now, 0, 0, 0)
sessions, err := store.TokenStatsBySession()
if err != nil {
t.Fatal(err)
}
if len(sessions) != 0 {
t.Fatalf("cancelled intent leaked into sessions: %+v", sessions)
}
var rollups, datedRollups int
if err := d.QueryRow(`SELECT COUNT(*), COUNT(*) FILTER (
WHERE metadata->>'token_day'=TO_CHAR(created_at AT TIME ZONE 'UTC', 'YYYY-MM-DD')
) FROM activity
WHERE exploration_id=$1 AND worker='token-ledger' AND kind='result'
AND metadata->>'cancelled_intent_id'=$2`, expID, fmt.Sprint(intentID)).Scan(&rollups, &datedRollups); err != nil {
t.Fatal(err)
}
if rollups != 3 || datedRollups != rollups {
t.Fatalf("token rollups=%d dated=%d, want three correctly dated rows", rollups, datedRollups)
}
if _, err := store.CancelIntent(intentID); err == nil {
t.Fatal("second cancellation unexpectedly succeeded")
}
afterRetry, err := store.TokenTotal()
if err != nil {
t.Fatal(err)
}
assertTokenUsage(t, afterRetry, 60, 63, 66, 69)
afterRetryDaily, err := d.TokenDailyAll(30)
if err != nil {
t.Fatal(err)
}
assertDailyTokenBucketsEqual(t, dailyTokenBuckets(afterRetryDaily), afterCancel)
}
func dailyTokenBuckets(items []DailyTokenBucket) map[string]DailyTokenBucket {
out := make(map[string]DailyTokenBucket, len(items))
for _, item := range items {
out[item.Day] = item
}
return out
}
func assertDailyTokenDelta(t *testing.T, got, baseline map[string]DailyTokenBucket, at time.Time, input, output, read int) {
t.Helper()
day := at.UTC().Format(time.DateOnly)
actual, initial := got[day], baseline[day]
if actual.InputTokens-initial.InputTokens != input ||
actual.OutputTokens-initial.OutputTokens != output ||
actual.CacheReadTokens-initial.CacheReadTokens != read {
t.Fatalf("token delta for %s = input:%d output:%d cache-read:%d, want %d/%d/%d",
day, actual.InputTokens-initial.InputTokens, actual.OutputTokens-initial.OutputTokens,
actual.CacheReadTokens-initial.CacheReadTokens, input, output, read)
}
}
func assertDailyTokenBucketsEqual(t *testing.T, got, want map[string]DailyTokenBucket) {
t.Helper()
if len(got) != len(want) {
t.Fatalf("daily token buckets changed after retry: got=%+v want=%+v", got, want)
}
for day, expected := range want {
if actual, ok := got[day]; !ok || actual != expected {
t.Fatalf("daily token bucket %s changed after retry: got=%+v want=%+v", day, actual, expected)
}
}
}
func TestCancelIntentPreservesOutputsYieldedByAnotherIntent(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
expID, err := d.CreateExploration("shared intent output", "preserve shared facts and findings")
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM explorations WHERE id=$1`, expID)
store := d.Exploration(expID)
first, err := store.AddIntent(map[string]any{"summary": "first"}, 1, nil, "planner")
if err != nil {
t.Fatal(err)
}
second, err := store.AddIntent(map[string]any{"summary": "second"}, 1, nil, "planner")
if err != nil {
t.Fatal(err)
}
fact, err := store.AddNode(KindFact, map[string]any{"text": "shared fact"}, 1, "confirmed", "worker", nil)
if err != nil {
t.Fatal(err)
}
finding, err := store.AddNode(KindFinding, map[string]any{"summary": "shared finding"}, 1, "confirmed", "worker", nil)
if err != nil {
t.Fatal(err)
}
for _, intentID := range []int64{first, second} {
if err := store.Link(intentID, RelYields, fact); err != nil {
t.Fatal(err)
}
if err := store.Link(intentID, RelYields, finding); err != nil {
t.Fatal(err)
}
}
if claimed, err := store.ClaimIntent(first, "work#1"); err != nil || !claimed {
t.Fatalf("claim: claimed=%v err=%v", claimed, err)
}
if err := store.SetIntentState(first, "paused"); err != nil {
t.Fatal(err)
}
cleanup, err := store.CancelIntent(first)
if err != nil {
t.Fatal(err)
}
if cleanup.Intents != 1 || cleanup.Facts != 0 || cleanup.Findings != 0 {
t.Fatalf("cleanup=%+v, want only the cancelled intent", cleanup)
}
for _, nodeID := range []int64{fact, finding} {
node, getErr := store.GetNode(nodeID)
if getErr != nil || node == nil {
t.Fatalf("shared node %d was removed: node=%+v err=%v", nodeID, node, getErr)
}
}
}
+170
View File
@@ -0,0 +1,170 @@
package db
import "testing"
// mustIntent / mustNode / mustLink are terse builders for the delete-cascade tests.
func mustIntent(t *testing.T, es *ExplorationStore, summary string) int64 {
t.Helper()
id, err := es.AddIntent(map[string]any{"summary": summary}, 1, nil, "planner")
if err != nil {
t.Fatal(err)
}
return id
}
func mustNode(t *testing.T, es *ExplorationStore, kind, summary string) int64 {
t.Helper()
id, err := es.AddNode(kind, map[string]any{"summary": summary}, 1, "confirmed", "worker", nil)
if err != nil {
t.Fatal(err)
}
return id
}
func mustLink(t *testing.T, es *ExplorationStore, from int64, rel string, to int64) {
t.Helper()
if err := es.Link(from, rel, to); err != nil {
t.Fatal(err)
}
}
func gone(t *testing.T, es *ExplorationStore, id int64) bool {
t.Helper()
n, err := es.GetNode(id)
if err != nil {
t.Fatal(err)
}
return n == nil
}
// TestSoftDeleteIntent 假删除置 deleted + delete_reason,保留节点。
func TestSoftDeleteIntent(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
expID, err := d.CreateExploration("soft delete", "假删除")
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM explorations WHERE id=$1`, expID)
es := d.Exploration(expID)
intent := mustIntent(t, es, "待删意图")
if err := es.SetIntentState(intent, "paused"); err != nil {
t.Fatal(err)
}
summary, err := es.SoftDeleteIntent(intent, "方向判断错误")
if err != nil {
t.Fatal(err)
}
if summary != "待删意图" {
t.Fatalf("summary=%q, want 待删意图", summary)
}
n, err := es.GetNode(intent)
if err != nil || n == nil {
t.Fatalf("intent removed by soft delete: n=%+v err=%v", n, err)
}
if n.State != StateIntentDeleted || n.DeleteReason != "方向判断错误" {
t.Fatalf("state=%q delete_reason=%q, want deleted/方向判断错误", n.State, n.DeleteReason)
}
// 待领(open)意图也允许假删除。
openIntent := mustIntent(t, es, "待领意图")
if _, err := es.SoftDeleteIntent(openIntent, "方向不需要了"); err != nil {
t.Fatalf("soft delete open intent: %v", err)
}
if n, err := es.GetNode(openIntent); err != nil || n == nil || n.State != StateIntentDeleted {
t.Fatalf("open intent not soft-deleted: n=%+v err=%v", n, err)
}
// 已删除(deleted)等其它状态不能再假删除。
if _, err := es.SoftDeleteIntent(intent, "再删"); err == nil {
t.Fatal("soft-deleting an already-deleted intent unexpectedly succeeded")
}
}
// TestHardDeleteCascadesExclusiveDescendants 真删除沿链路级联删除独占子孙到叶子。
func TestHardDeleteCascadesExclusiveDescendants(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
expID, err := d.CreateExploration("hard cascade", "级联删除")
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM explorations WHERE id=$1`, expID)
es := d.Exploration(expID)
// intent1 --yields--> fact1 --derived_from--> intent2 --yields--> fact2(叶子)
intent1 := mustIntent(t, es, "根意图")
fact1 := mustNode(t, es, KindFact, "事实1")
mustLink(t, es, intent1, RelYields, fact1)
intent2 := mustIntent(t, es, "衍生意图")
mustLink(t, es, fact1, RelDerivedFrom, intent2)
fact2 := mustNode(t, es, KindFact, "事实2")
mustLink(t, es, intent2, RelYields, fact2)
cleanup, err := es.CancelIntent(intent1)
if err != nil {
t.Fatal(err)
}
if cleanup.Intents != 2 || cleanup.Facts != 2 {
t.Fatalf("cleanup=%+v, want 2 intents / 2 facts", cleanup)
}
for _, id := range []int64{intent1, fact1, intent2, fact2} {
if !gone(t, es, id) {
t.Fatalf("node %d survived cascade", id)
}
}
}
// TestHardDeletePreservesSharedAndGoal 真删除保留共享子孙(还有其它父)与目标。
func TestHardDeletePreservesSharedAndGoal(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
expID, err := d.CreateExploration("hard preserve", "保留共享/目标")
if err != nil {
t.Fatal(err)
}
defer d.Exec(`DELETE FROM explorations WHERE id=$1`, expID)
es := d.Exploration(expID)
goal, err := es.AddGoal(map[string]any{"text": "拿下后台"}, "human")
if err != nil {
t.Fatal(err)
}
// intent1 独占 finding(proves goal),intent1 与 intentX 共享 fact1(fact1 衍生出 intent2)。
intent1 := mustIntent(t, es, "待删意图")
intentX := mustIntent(t, es, "旁路意图")
finding := mustNode(t, es, KindFinding, "漏洞")
mustLink(t, es, intent1, RelYields, finding)
mustLink(t, es, finding, RelProves, goal)
shared := mustNode(t, es, KindFact, "共享事实")
mustLink(t, es, intent1, RelYields, shared)
mustLink(t, es, intentX, RelYields, shared)
intent2 := mustIntent(t, es, "由共享事实衍生")
mustLink(t, es, shared, RelDerivedFrom, intent2)
cleanup, err := es.CancelIntent(intent1)
if err != nil {
t.Fatal(err)
}
// 只删 intent1 与其独占的 finding;shared(有 intentX 父)及其下游 intent2、goal 全保留。
if cleanup.Intents != 1 || cleanup.Findings != 1 || cleanup.Facts != 0 {
t.Fatalf("cleanup=%+v, want 1 intent / 1 finding / 0 fact", cleanup)
}
if !gone(t, es, intent1) || !gone(t, es, finding) {
t.Fatal("intent1/finding should be removed")
}
for _, id := range []int64{goal, intentX, shared, intent2} {
if gone(t, es, id) {
t.Fatalf("node %d was wrongly cascaded", id)
}
}
}
+374
View File
@@ -0,0 +1,374 @@
package db
import (
"database/sql"
"encoding/json"
"fmt"
"strings"
"time"
)
// InterceptRule is one row of intercept_rules.
type InterceptRule struct {
ID int64 `json:"id"`
Name string `json:"name"`
Enabled bool `json:"enabled"`
Priority int `json:"priority"`
MatchTarget string `json:"match_target"`
MatchType string `json:"match_type"`
Pattern string `json:"pattern"`
Action string `json:"action"`
Message string `json:"message"`
TimeoutEnabled bool `json:"timeout_enabled"`
TimeoutSeconds int `json:"timeout_seconds"`
TimeoutAction string `json:"timeout_action"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// InterceptPending is one row of intercept_pending.
type InterceptPending struct {
ID int64 `json:"id"`
RuleID *int64 `json:"rule_id"`
ConversationID *int64 `json:"conversation_id"`
TaskID *string `json:"task_id"`
AgentName string `json:"agent_name"`
ToolName string `json:"tool_name"`
ToolInput json.RawMessage `json:"tool_input"`
Status string `json:"status"`
DecisionSource string `json:"decision_source"`
Reason string `json:"reason"` // 规则 message 或模型判定理由(前缀 [模型])
DecidedAt *time.Time `json:"decided_at"`
CreatedAt time.Time `json:"created_at"`
}
const interceptRuleCols = `id, name, enabled, priority, match_target, match_type, pattern, action, message, timeout_enabled, timeout_seconds, timeout_action, created_at, updated_at`
func scanInterceptRule(row interface{ Scan(...any) error }) (InterceptRule, error) {
var r InterceptRule
err := row.Scan(&r.ID, &r.Name, &r.Enabled, &r.Priority,
&r.MatchTarget, &r.MatchType, &r.Pattern, &r.Action, &r.Message,
&r.TimeoutEnabled, &r.TimeoutSeconds, &r.TimeoutAction,
&r.CreatedAt, &r.UpdatedAt)
return r, err
}
// ListInterceptRules returns all rules ordered by priority DESC then id.
func (d *DB) ListInterceptRules() ([]InterceptRule, error) {
rows, err := d.Query(`SELECT ` + interceptRuleCols + ` FROM intercept_rules ORDER BY priority DESC, id`)
if err != nil {
return nil, err
}
defer rows.Close()
var out []InterceptRule
for rows.Next() {
r, err := scanInterceptRule(rows)
if err != nil {
return nil, err
}
out = append(out, r)
}
return out, rows.Err()
}
// CreateInterceptRule inserts a new rule.
func (d *DB) CreateInterceptRule(name, matchTarget, matchType, pattern, action, message string, priority int, enabled bool, timeoutEnabled bool, timeoutSeconds int, timeoutAction string) (InterceptRule, error) {
row := d.QueryRow(`
INSERT INTO intercept_rules(name, enabled, priority, match_target, match_type, pattern, action, message, timeout_enabled, timeout_seconds, timeout_action)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11)
RETURNING `+interceptRuleCols,
name, enabled, priority, matchTarget, matchType, pattern, action, message, timeoutEnabled, timeoutSeconds, timeoutAction)
return scanInterceptRule(row)
}
// UpdateInterceptRule replaces all editable fields of an existing rule.
func (d *DB) UpdateInterceptRule(id int64, name, matchTarget, matchType, pattern, action, message string, priority int, enabled bool, timeoutEnabled bool, timeoutSeconds int, timeoutAction string) (InterceptRule, error) {
row := d.QueryRow(`
UPDATE intercept_rules
SET name=$2, enabled=$3, priority=$4, match_target=$5,
match_type=$6, pattern=$7, action=$8, message=$9,
timeout_enabled=$10, timeout_seconds=$11, timeout_action=$12
WHERE id=$1
RETURNING `+interceptRuleCols,
id, name, enabled, priority, matchTarget, matchType, pattern, action, message, timeoutEnabled, timeoutSeconds, timeoutAction)
return scanInterceptRule(row)
}
// DeleteInterceptRule removes a rule.
func (d *DB) DeleteInterceptRule(id int64) error {
_, err := d.Exec(`DELETE FROM intercept_rules WHERE id=$1`, id)
return err
}
// ToggleInterceptRule flips the enabled state of a rule.
func (d *DB) ToggleInterceptRule(id int64, enabled bool) error {
_, err := d.Exec(`UPDATE intercept_rules SET enabled=$2 WHERE id=$1`, id, enabled)
return err
}
// CreateInterceptPending inserts a pending approval record and returns its ID.
// convID == 0 → conversation_id stored as NULL (background task).
// taskID == "" → task_id stored as NULL.
func (d *DB) CreateInterceptPending(ruleID, convID int64, taskID, agentName, toolName string, input []byte, reason string, audits ...*InterceptAudit) (int64, error) {
raw := json.RawMessage(input)
if len(raw) == 0 {
raw = json.RawMessage("{}")
}
var convIDPtr *int64
if convID != 0 {
convIDPtr = &convID
}
var taskIDPtr *string
if taskID != "" {
taskIDPtr = &taskID
}
// ruleID == 0 → NULL: the LLM fallback judge has no owning rule.
var ruleIDPtr *int64
if ruleID != 0 {
ruleIDPtr = &ruleID
}
var id int64
err := d.QueryRow(`
INSERT INTO intercept_pending(rule_id, conversation_id, task_id, agent_name, tool_name, tool_input, reason, decision_source, audit)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9) RETURNING id`,
ruleIDPtr, convIDPtr, taskIDPtr, agentName, toolName, raw, reason, interceptSource(ruleID, reason), firstAudit(audits)).Scan(&id)
return id, err
}
// DecideInterceptPending updates a pending record's status (allowed/denied/timeout).
func (d *DB) DecideInterceptPending(id int64, status string) error {
_, err := d.Exec(`UPDATE intercept_pending SET status=$2, decided_at=NOW() WHERE id=$1`, id, status)
return err
}
// CreateDecidedIntercept inserts an intercept_pending row ALREADY in a final state
// (status = 'allowed' | 'denied'), decided_at stamped now. Used to log allow/deny
// rule matches for observability — they don't block and need no user action, so unlike
// CreateInterceptPending (which starts 'pending') this records the outcome directly.
func (d *DB) CreateDecidedIntercept(ruleID, convID int64, taskID, agentName, toolName string, input []byte, status, reason string, audits ...*InterceptAudit) (int64, error) {
raw := json.RawMessage(input)
if len(raw) == 0 {
raw = json.RawMessage("{}")
}
var convIDPtr *int64
if convID != 0 {
convIDPtr = &convID
}
var taskIDPtr *string
if taskID != "" {
taskIDPtr = &taskID
}
// ruleID == 0 → NULL: the LLM fallback judge has no owning rule.
var ruleIDPtr *int64
if ruleID != 0 {
ruleIDPtr = &ruleID
}
var id int64
err := d.QueryRow(`
INSERT INTO intercept_pending(rule_id, conversation_id, task_id, agent_name, tool_name, tool_input, status, reason, decided_at, decision_source, audit)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, NOW(), $9, $10) RETURNING id`,
ruleIDPtr, convIDPtr, taskIDPtr, agentName, toolName, raw, status, reason, interceptSource(ruleID, reason), firstAudit(audits)).Scan(&id)
return id, err
}
const interceptPendingCols = `id, rule_id, conversation_id, task_id, agent_name, tool_name, tool_input, status, reason, decided_at, created_at, decision_source`
func scanInterceptPending(s interface{ Scan(...any) error }, p *InterceptPending) error {
return s.Scan(&p.ID, &p.RuleID, &p.ConversationID, &p.TaskID, &p.AgentName,
&p.ToolName, &p.ToolInput, &p.Status, &p.Reason, &p.DecidedAt, &p.CreatedAt, &p.DecisionSource)
}
// ListPendingIntercepts returns all unresolved approval requests, newest first.
func (d *DB) ListPendingIntercepts() ([]InterceptPending, error) {
rows, err := d.Query(`SELECT ` + interceptPendingCols + ` FROM intercept_pending WHERE status='pending' ORDER BY created_at DESC`)
if err != nil {
return nil, err
}
defer rows.Close()
var out []InterceptPending
for rows.Next() {
var p InterceptPending
if err := scanInterceptPending(rows, &p); err != nil {
return nil, err
}
out = append(out, p)
}
return out, rows.Err()
}
// GetInterceptPending returns one pending record (nil if absent).
func (d *DB) GetInterceptPending(id int64) (*InterceptPending, error) {
var p InterceptPending
err := scanInterceptPending(
d.QueryRow(`SELECT `+interceptPendingCols+` FROM intercept_pending WHERE id=$1`, id),
&p,
)
if err == sql.ErrNoRows {
return nil, nil
}
return &p, err
}
// InterceptApprovalRow is intercept_pending enriched with conversation and rule info.
type InterceptApprovalRow struct {
InterceptPending
ConvTitle string `json:"conv_title"`
ConvAgentKey string `json:"conv_agent_key"`
RuleName string `json:"rule_name"`
}
func scanInterceptApprovalRow(rows interface{ Scan(...any) error }, r *InterceptApprovalRow) error {
return rows.Scan(
&r.ID, &r.RuleID, &r.ConversationID, &r.TaskID, &r.AgentName,
&r.ToolName, &r.ToolInput, &r.Status, &r.Reason, &r.DecidedAt, &r.CreatedAt,
&r.DecisionSource, &r.ConvTitle, &r.ConvAgentKey, &r.RuleName,
)
}
// Keep legacy rows without decision_source consistent with their displayed source.
const approvalDecisionSource = `COALESCE(NULLIF(ip.decision_source,''), CASE
WHEN ip.rule_id IS NOT NULL THEN 'rule'
WHEN ip.reason LIKE '[模型]%' THEN 'model' ELSE 'unknown' END)`
const approvalRowColumns = `ip.id, ip.rule_id, ip.conversation_id, ip.task_id, ip.agent_name,
ip.tool_name, ip.tool_input, ip.status, ip.reason, ip.decided_at, ip.created_at, ` + approvalDecisionSource + `,
COALESCE(c.title,'') AS conv_title,
COALESCE(c.agent_key,'') AS conv_agent_key,
COALESCE(ir.name,'') AS rule_name`
const approvalRowJoins = `
FROM intercept_pending ip
LEFT JOIN conversations c ON c.id = ip.conversation_id
LEFT JOIN intercept_rules ir ON ir.id = ip.rule_id`
const approvalRowSelect = `SELECT ` + approvalRowColumns + approvalRowJoins
const approvalRowSelectWithAudit = `SELECT ` + approvalRowColumns + `, ip.audit` + approvalRowJoins
func interceptSource(ruleID int64, reason string) string {
if ruleID != 0 {
return "rule"
}
if strings.HasPrefix(reason, "[模型]") {
return "model"
}
return "unknown"
}
func firstAudit(audits []*InterceptAudit) any {
if len(audits) == 0 || audits[0] == nil {
return nil
}
raw, err := json.Marshal(audits[0])
if err != nil {
return nil
}
return raw
}
// ListAllIntercepts returns up to limit intercept_pending rows (newest first)
// joined with conversation and rule info.
func (d *DB) ListAllIntercepts(limit int) ([]InterceptApprovalRow, error) {
rows, err := d.Query(approvalRowSelect+` ORDER BY ip.created_at DESC LIMIT $1`, limit)
if err != nil {
return nil, err
}
defer rows.Close()
var out []InterceptApprovalRow
for rows.Next() {
var r InterceptApprovalRow
if err := scanInterceptApprovalRow(rows, &r); err != nil {
return nil, err
}
out = append(out, r)
}
return out, rows.Err()
}
// InterceptApprovalFilter combines exact status and decision-source filters.
// Empty fields include all values.
type InterceptApprovalFilter struct {
Status string
DecisionSource string
}
// ListAllInterceptsPage returns one 1-based page and the total matching count.
func (d *DB) ListAllInterceptsPage(page, size int, filter InterceptApprovalFilter) ([]InterceptApprovalRow, int, error) {
return d.listInterceptsPage("", page, size, filter)
}
// ListTaskIntercepts returns all intercept_pending rows for a specific task (newest first).
func (d *DB) ListTaskIntercepts(taskID string) ([]InterceptApprovalRow, error) {
rows, err := d.Query(approvalRowSelect+` WHERE ip.task_id=$1 ORDER BY ip.created_at DESC`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
var out []InterceptApprovalRow
for rows.Next() {
var r InterceptApprovalRow
if err := scanInterceptApprovalRow(rows, &r); err != nil {
return nil, err
}
out = append(out, r)
}
return out, rows.Err()
}
// ListTaskInterceptsPage is the paginated variant of ListTaskIntercepts.
func (d *DB) ListTaskInterceptsPage(taskID string, page, size int, filter InterceptApprovalFilter) ([]InterceptApprovalRow, int, error) {
return d.listInterceptsPage(taskID, page, size, filter)
}
func (d *DB) listInterceptsPage(taskID string, page, size int, filter InterceptApprovalFilter) ([]InterceptApprovalRow, int, error) {
if page < 1 {
page = 1
}
if size <= 0 {
size = 20
}
if size > 100 {
size = 100
}
offset := (page - 1) * size
conditions := []string{}
args := []any{}
add := func(column, value string) {
if value != "" {
args = append(args, value)
conditions = append(conditions, column+"=$"+fmt.Sprint(len(args)))
}
}
add("ip.task_id", taskID)
add("ip.status", filter.Status)
add(approvalDecisionSource, filter.DecisionSource)
where := ""
if len(conditions) > 0 {
where = " WHERE " + strings.Join(conditions, " AND ")
}
var total int
if err := d.QueryRow("SELECT COUNT(*) FROM intercept_pending ip"+where, args...).Scan(&total); err != nil {
return nil, 0, err
}
limitArg := len(args) + 1
offsetArg := limitArg + 1
dataQ := approvalRowSelect + where +
" ORDER BY ip.created_at DESC, ip.id DESC LIMIT $" + fmt.Sprint(limitArg) +
" OFFSET $" + fmt.Sprint(offsetArg)
args = append(args, size, offset)
rows, err := d.Query(dataQ, args...)
if err != nil {
return nil, 0, err
}
defer rows.Close()
out := []InterceptApprovalRow{}
for rows.Next() {
var r InterceptApprovalRow
if err := scanInterceptApprovalRow(rows, &r); err != nil {
return nil, 0, err
}
out = append(out, r)
}
return out, total, rows.Err()
}
+114
View File
@@ -0,0 +1,114 @@
package db
import (
"database/sql"
"encoding/json"
"time"
)
// InterceptContextEntry is a bounded, recorded session event, not model reasoning.
type InterceptContextEntry struct {
Kind string `json:"kind"`
Tool string `json:"tool,omitempty"`
ToolUseID string `json:"tool_use_id,omitempty"`
Text string `json:"text"`
IsError bool `json:"is_error,omitempty"`
Truncated bool `json:"truncated,omitempty"`
}
// InterceptAudit is captured at review time. It is deliberately excluded from
// polling/list responses; old rows have no audit instead of reconstructed data.
type InterceptAudit struct {
RunID string `json:"run_id,omitempty"`
ToolUseID string `json:"tool_use_id,omitempty"`
Correlation string `json:"correlation"` // exact | ambiguous | unavailable
InputDigest string `json:"input_digest"`
UserMessage string `json:"user_message"`
UserTruncated bool `json:"user_truncated,omitempty"`
Context []InterceptContextEntry `json:"context"`
ContextTruncated bool `json:"context_truncated,omitempty"`
CapturedAt time.Time `json:"captured_at"`
ModelFallback bool `json:"model_fallback,omitempty"`
ModelInput json.RawMessage `json:"model_input,omitempty"`
ModelInputDigest string `json:"model_input_digest,omitempty"`
InitialAction string `json:"initial_action"`
InitialReason string `json:"initial_reason"`
EffectiveAction string `json:"effective_action,omitempty"`
DecisionReason string `json:"decision_reason,omitempty"`
RuleName string `json:"rule_name,omitempty"`
ConfigDigest string `json:"config_digest,omitempty"`
ProfileID int64 `json:"profile_id,omitempty"`
ExecutionStatus string `json:"execution_status"`
Output string `json:"output,omitempty"`
OutputTruncated bool `json:"output_truncated,omitempty"`
ExecutionEndedAt *time.Time `json:"execution_ended_at,omitempty"`
}
type InterceptDetail struct {
InterceptApprovalRow
Audit *InterceptAudit `json:"audit"`
}
func (d *DB) GetInterceptDetail(id int64) (*InterceptDetail, error) {
var out InterceptDetail
var raw []byte
err := d.QueryRow(approvalRowSelectWithAudit+` WHERE ip.id=$1`, id).Scan(
&out.ID, &out.RuleID, &out.ConversationID, &out.TaskID, &out.AgentName,
&out.ToolName, &out.ToolInput, &out.Status, &out.Reason, &out.DecidedAt, &out.CreatedAt,
&out.DecisionSource, &out.ConvTitle, &out.ConvAgentKey, &out.RuleName, &raw,
)
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
if len(raw) > 0 {
if err := json.Unmarshal(raw, &out.Audit); err != nil {
return nil, err
}
}
return &out, nil
}
// ResolveIntercept atomically settles a pending request. A timeout cannot
// overwrite a human decision and repeat decisions cannot rewrite history.
func (d *DB) ResolveIntercept(id int64, status, action, reason string) (bool, error) {
execution := "not_executed"
if action == "allow" {
execution = "awaiting_result"
}
patch, err := json.Marshal(map[string]any{
"effective_action": action, "decision_reason": reason, "execution_status": execution,
})
if err != nil {
return false, err
}
r, err := d.Exec(`UPDATE intercept_pending SET status=$2, decided_at=NOW(),
audit=CASE WHEN audit IS NULL THEN NULL ELSE audit || $3::jsonb ||
CASE WHEN $4='allow' AND audit->>'correlation' IS DISTINCT FROM 'exact'
THEN '{"execution_status":"unknown"}'::jsonb ELSE '{}'::jsonb END END
WHERE id=$1 AND status='pending'`, id, status, patch, action)
if err != nil {
return false, err
}
n, err := r.RowsAffected()
return n == 1, err
}
// CompleteIntercept only updates the exact recorded call after it was allowed.
// A blocked tool_result must never be presented as a failed execution.
func (d *DB) CompleteIntercept(id int64, runID, toolUseID, status, output string, truncated bool) error {
patch, err := json.Marshal(map[string]any{
"execution_status": status, "output": output, "output_truncated": truncated,
"execution_ended_at": time.Now().UTC(),
})
if err != nil {
return err
}
_, err = d.Exec(`UPDATE intercept_pending SET audit=audit || $4::jsonb
WHERE id=$1 AND audit->>'run_id'=$2 AND audit->>'tool_use_id'=$3
AND audit->>'effective_action'='allow' AND audit->>'execution_status'='awaiting_result'`,
id, runID, toolUseID, patch)
return err
}
+167
View File
@@ -0,0 +1,167 @@
package db
import (
"encoding/json"
"strings"
"sync"
"sync/atomic"
"testing"
)
func TestInterceptDetails(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.Close() })
create := func(t *testing.T, audit *InterceptAudit) int64 {
t.Helper()
id, err := d.CreateInterceptPending(0, 0, "approval-detail-test", "test", "Write", []byte(`{"path":"report.md"}`), "[模型] 请确认", audit)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _, _ = d.Exec(`DELETE FROM intercept_pending WHERE id=$1`, id) })
return id
}
t.Run("legacy and lazy payload", func(t *testing.T) {
id := create(t, nil)
got, err := d.GetInterceptDetail(id)
if err != nil || got == nil || got.Audit != nil {
t.Fatalf("legacy: %+v %v", got, err)
}
create(t, &InterceptAudit{UserMessage: "snapshot-only-marker", InitialAction: "ask"})
items, err := d.ListTaskIntercepts("approval-detail-test")
if err != nil {
t.Fatal(err)
}
raw, _ := json.Marshal(items)
if strings.Contains(string(raw), "snapshot-only-marker") || strings.Contains(string(raw), `"audit"`) {
t.Fatal("snapshot leaked into list response")
}
missing, err := d.GetInterceptDetail(-1)
if err != nil || missing != nil {
t.Fatal("missing record not reported")
}
})
t.Run("decision race and exact output", func(t *testing.T) {
id := create(t, &InterceptAudit{RunID: "run", ToolUseID: "call", Correlation: "exact", InitialAction: "ask", ExecutionStatus: "not_started"})
var wins atomic.Int32
var wg sync.WaitGroup
for range 8 {
wg.Go(func() {
ok, err := d.ResolveIntercept(id, "allowed", "allow", "人工允许执行")
if err != nil {
t.Error(err)
}
if ok {
wins.Add(1)
}
})
}
wg.Wait()
if wins.Load() != 1 {
t.Fatalf("%d decisions won", wins.Load())
}
ok, err := d.ResolveIntercept(id, "timeout", "deny", "late timeout")
if err != nil || ok {
t.Fatal("timeout overwrote decision")
}
if err := d.CompleteIntercept(id, "different-run", "call", "succeeded", "WRONG", false); err != nil {
t.Fatal(err)
}
got, _ := d.GetInterceptDetail(id)
if got.Audit.Output != "" {
t.Fatal("cross-run result attached")
}
if err := d.CompleteIntercept(id, "run", "call", "failed", "permission denied", true); err != nil {
t.Fatal(err)
}
got, err = d.GetInterceptDetail(id)
if err != nil || got.Status != "allowed" || got.Audit.InitialAction != "ask" || got.Audit.ExecutionStatus != "failed" || !got.Audit.OutputTruncated {
t.Fatalf("wrong details: %+v %v", got, err)
}
if err := d.CompleteIntercept(id, "run", "call", "succeeded", "late duplicate", false); err != nil {
t.Fatal(err)
}
got, _ = d.GetInterceptDetail(id)
if got.Audit.Output != "permission denied" {
t.Fatal("duplicate result rewrote output")
}
})
t.Run("denied output and timeout allow", func(t *testing.T) {
id := create(t, &InterceptAudit{RunID: "run", ToolUseID: "call", ExecutionStatus: "not_started"})
if _, err := d.ResolveIntercept(id, "denied", "deny", "人工拒绝"); err != nil {
t.Fatal(err)
}
if err := d.CompleteIntercept(id, "run", "call", "failed", "Blocked by hook", false); err != nil {
t.Fatal(err)
}
got, _ := d.GetInterceptDetail(id)
if got.Audit.ExecutionStatus != "not_executed" || got.Audit.Output != "" {
t.Fatal("denial presented as executed")
}
id = create(t, &InterceptAudit{RunID: "run2", ToolUseID: "call2", Correlation: "exact", InitialAction: "ask"})
if _, err := d.ResolveIntercept(id, "timeout", "allow", "超时允许"); err != nil {
t.Fatal(err)
}
if err := d.CompleteIntercept(id, "run2", "call2", "succeeded", "ok", false); err != nil {
t.Fatal(err)
}
got, _ = d.GetInterceptDetail(id)
if got.Status != "timeout" || got.Audit.EffectiveAction != "allow" || got.Audit.ExecutionStatus != "succeeded" {
t.Fatal("timeout action lost")
}
})
t.Run("archive compatibility", func(t *testing.T) {
for _, legacy := range []bool{true, false} {
id := create(t, &InterceptAudit{InitialAction: "ask", ExecutionStatus: "not_started"})
var raw []byte
if err := d.QueryRow(`SELECT row_to_json(ip) FROM intercept_pending ip WHERE id=$1`, id).Scan(&raw); err != nil {
t.Fatal(err)
}
var row map[string]any
if err := json.Unmarshal(raw, &row); err != nil {
t.Fatal(err)
}
if legacy {
delete(row, "audit")
delete(row, "decision_source")
}
archived, _ := json.Marshal([]map[string]any{row})
tx, err := d.Begin()
if err != nil {
t.Fatal(err)
}
defer tx.Rollback()
if _, err := tx.Exec(`DELETE FROM intercept_pending WHERE id=$1`, id); err != nil {
t.Fatal(err)
}
if err := restoreInterceptRows(tx, archived); err != nil {
t.Fatalf("legacy=%t: %v", legacy, err)
}
var status, source string
var auditJSON []byte
if err := tx.QueryRow(`SELECT status, decision_source, audit FROM intercept_pending WHERE id=$1`, id).Scan(&status, &source, &auditJSON); err != nil {
t.Fatal(err)
}
if status != "timeout" || source != "model" {
t.Fatalf("restored %s/%s", status, source)
}
if legacy && len(auditJSON) > 0 {
t.Fatal("fabricated legacy audit")
}
if !legacy {
var a InterceptAudit
if err := json.Unmarshal(auditJSON, &a); err != nil {
t.Fatal(err)
}
if a.InitialAction != "ask" || a.EffectiveAction != "deny" || a.ExecutionStatus != "not_executed" {
t.Fatalf("restored audit: %+v", a)
}
}
if err := tx.Rollback(); err != nil {
t.Fatal(err)
}
}
})
}
+120
View File
@@ -0,0 +1,120 @@
package db
import (
"errors"
"fmt"
"strconv"
)
var ErrInterceptTaskDeleted = errors.New("任务已被删除或归档")
var ErrInterceptSessionDeleted = errors.New("对应会话或执行记录已被删除或不存在")
var ErrInterceptExecutionUnavailable = errors.New("未找到可唯一关联的原始工具调用;记录可能已删除,或旧审批没有保存关联 ID")
// InterceptExecution is a navigation target read from original activity rows.
// It is not model context and never falls back to matching command text.
type InterceptExecution struct {
ConversationID *int64 `json:"conversation_id,omitempty"`
TaskID *string `json:"task_id,omitempty"`
Session string `json:"session"`
Seq int64 `json:"seq"`
Items []Activity `json:"-"`
}
func (d *DB) GetInterceptExecution(id int64) (*InterceptExecution, error) {
approval, err := d.GetInterceptDetail(id)
if err != nil || approval == nil {
return nil, err
}
audit := approval.Audit
if audit == nil || audit.Correlation != "exact" || audit.ToolUseID == "" {
return nil, ErrInterceptExecutionUnavailable
}
var query string
var scope any
if approval.ConversationID != nil {
scope = *approval.ConversationID
query = `SELECT id, NULL::bigint, COALESCE(worker,''), kind, COALESCE(tool,''), tool_use_id, is_error, COALESCE(summary,''), created_at, NULL::integer FROM conversation_activities WHERE conversation_id=$1`
} else if approval.TaskID != nil {
taskID, parseErr := strconv.ParseInt(*approval.TaskID, 10, 64)
if parseErr != nil || taskID <= 0 {
return nil, ErrInterceptExecutionUnavailable
}
var exists bool
if err := d.QueryRow(`SELECT EXISTS(SELECT 1 FROM tasks WHERE id=$1 AND archived_at IS NULL AND deleted_at IS NULL)`, taskID).Scan(&exists); err != nil {
return nil, err
}
if !exists {
return nil, ErrInterceptTaskDeleted
}
scope = taskID
query = `SELECT id, node_id, COALESCE(worker,''), kind, COALESCE(tool,''), tool_use_id, is_error, COALESCE(summary,''), created_at, main_seg FROM activity WHERE exploration_id=(SELECT exploration_id FROM tasks WHERE id=$1 AND archived_at IS NULL AND deleted_at IS NULL)`
} else {
return nil, ErrInterceptExecutionUnavailable
}
// Three matches suffice to detect duplicate IDs without loading a transcript.
rows, err := d.Query(query+` AND tool_use_id=$2 AND kind IN ('tool_use','tool_result') ORDER BY id LIMIT 3`, scope, audit.ToolUseID)
if err != nil {
return nil, err
}
defer rows.Close()
out := &InterceptExecution{ConversationID: approval.ConversationID, TaskID: approval.TaskID}
for rows.Next() {
var a Activity
if err := rows.Scan(&a.ID, &a.NodeID, &a.Worker, &a.Kind, &a.Tool, &a.ToolUseID, &a.IsError, &a.Summary, &a.CreatedAt, &a.MainSeg); err != nil {
return nil, err
}
out.Items = append(out.Items, a)
}
if err := rows.Err(); err != nil {
return nil, err
}
if len(out.Items) == 0 {
return nil, ErrInterceptSessionDeleted
}
if len(out.Items) > 2 {
return nil, ErrInterceptExecutionUnavailable
}
call := out.Items[0]
if call.Kind != "tool_use" || call.Tool != approval.ToolName {
return nil, ErrInterceptExecutionUnavailable
}
if len(out.Items) == 2 {
result := out.Items[1]
if result.Kind != "tool_result" || result.Worker != call.Worker || (result.Tool != "" && result.Tool != call.Tool) || !sameOptionalInt64(call.NodeID, result.NodeID) || !sameMainSegment(call.MainSeg, result.MainSeg) {
return nil, ErrInterceptExecutionUnavailable
}
}
out.Seq = call.ID
if approval.ConversationID == nil {
switch {
case call.Worker == "mainagent":
seg := 0
if call.MainSeg != nil {
seg = *call.MainSeg
}
out.Session = fmt.Sprintf("main:%d", seg)
case call.Worker == "planner":
out.Session = "plan"
case call.NodeID != nil:
out.Session = fmt.Sprintf("intent:%d", *call.NodeID)
default:
return nil, ErrInterceptExecutionUnavailable
}
}
return out, nil
}
func sameOptionalInt64(a, b *int64) bool {
return (a == nil && b == nil) || (a != nil && b != nil && *a == *b)
}
func sameMainSegment(a, b *int) bool {
av, bv := 0, 0
if a != nil {
av = *a
}
if b != nil {
bv = *b
}
return av == bv
}
+150
View File
@@ -0,0 +1,150 @@
package db
import (
"errors"
"fmt"
"testing"
)
func TestInterceptExecutionNavigation(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.Close() })
approval := func(conv int64, task, id string, exact bool) int64 {
t.Helper()
audit := &InterceptAudit{ToolUseID: id, Correlation: "exact"}
if !exact {
audit.Correlation = "ambiguous"
}
n, err := d.CreateInterceptPending(0, conv, task, "display-name-is-not-a-session-id", "Bash", []byte(`{"command":"pwd"}`), "review", audit)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _, _ = d.Exec(`DELETE FROM intercept_pending WHERE id=$1`, n) })
return n
}
conv := func() int64 {
t.Helper()
c, err := d.CreateConversation("test", "navigation", nil)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _, _ = d.Exec(`DELETE FROM conversations WHERE id=$1`, c.ID) })
return c.ID
}
addConv := func(c int64, kind, id string) int64 {
t.Helper()
seq, err := d.AppendConvActivity(c, Activity{Worker: "test", Kind: kind, Tool: "Bash", ToolUseID: id, Summary: "pwd"})
if err != nil {
t.Fatal(err)
}
return seq
}
t.Run("conversation scope and old paginated call", func(t *testing.T) {
c1, c2 := conv(), conv()
seq := addConv(c1, "tool_use", "same-id")
addConv(c1, "tool_result", "same-id")
addConv(c2, "tool_use", "same-id")
addConv(c2, "tool_result", "same-id")
for range 210 {
addConv(c1, "text", "")
}
id := approval(c1, "", "same-id", true)
got, err := d.GetInterceptExecution(id)
if err != nil || got.Seq != seq || len(got.Items) != 2 || *got.ConversationID != c1 {
t.Fatalf("wrong conversation target: %+v %v", got, err)
}
addConv(c1, "tool_use", "same-id")
if _, err := d.GetInterceptExecution(id); !errors.Is(err, ErrInterceptExecutionUnavailable) {
t.Fatal("duplicate call ID selected an arbitrary execution")
}
})
t.Run("task sessions and result pairing", func(t *testing.T) {
task, err := d.CreateTask("navigation", "fixture", nil, 0, 600)
if err != nil {
t.Fatal(err)
}
exp := task.ExplorationID
t.Cleanup(func() {
_, _ = d.Exec(`DELETE FROM tasks WHERE id=$1`, task.ID)
_, _ = d.Exec(`DELETE FROM explorations WHERE id=$1`, exp)
})
es := d.Exploration(exp)
intent, err := es.AddIntent(map[string]any{"summary": "do not parse agent label"}, 1, nil, "planner")
if err != nil {
t.Fatal(err)
}
seg := 3
for _, tc := range []struct {
worker, key string
node *int64
segment *int
}{{"work#7", fmt.Sprintf("intent:%d", intent), &intent, nil}, {"planner", "plan", nil, nil}, {"mainagent", "main:3", nil, &seg}, {"mainagent", "main:0", nil, nil}} {
callID := "call-" + tc.key
a := Activity{Worker: tc.worker, NodeID: tc.node, MainSeg: tc.segment, Kind: "tool_use", Tool: "Bash", ToolUseID: callID, Summary: "pwd"}
seq, err := es.AppendActivity(a)
if err != nil {
t.Fatal(err)
}
a.Kind = "tool_result"
if _, err := es.AppendActivity(a); err != nil {
t.Fatal(err)
}
got, err := d.GetInterceptExecution(approval(0, fmt.Sprint(task.ID), callID, true))
if err != nil || got.Seq != seq || got.Session != tc.key || len(got.Items) != 2 {
t.Fatalf("wrong task session: %+v %v", got, err)
}
}
deletedCall := "deleted-session-call"
deletedSeq, err := es.AppendActivity(Activity{Worker: "work#9", NodeID: &intent, Kind: "tool_use", Tool: "Bash", ToolUseID: deletedCall})
if err != nil {
t.Fatal(err)
}
deletedApproval := approval(0, fmt.Sprint(task.ID), deletedCall, true)
if _, err = d.Exec(`DELETE FROM activity WHERE id=$1`, deletedSeq); err != nil {
t.Fatal(err)
}
if _, err = d.GetInterceptExecution(deletedApproval); !errors.Is(err, ErrInterceptSessionDeleted) {
t.Fatalf("deleted session: %v", err)
}
if _, err = d.Exec(`UPDATE tasks SET archived_at=NOW() WHERE id=$1`, task.ID); err != nil {
t.Fatal(err)
}
if _, err = d.GetInterceptExecution(deletedApproval); !errors.Is(err, ErrInterceptTaskDeleted) {
t.Fatalf("archived task: %v", err)
}
if _, err = d.Exec(`UPDATE tasks SET archived_at=NULL WHERE id=$1`, task.ID); err != nil {
t.Fatal(err)
}
a := Activity{Worker: "work#7", NodeID: &intent, Kind: "tool_use", Tool: "Bash", ToolUseID: "unpaired", Summary: "pwd"}
if _, err := es.AppendActivity(a); err != nil {
t.Fatal(err)
}
id := approval(0, fmt.Sprint(task.ID), "unpaired", true)
got, err := d.GetInterceptExecution(id)
if err != nil || len(got.Items) != 1 {
t.Fatal("pending call is not navigable")
}
a.Kind = "tool_result"
a.Worker = "different-worker"
if _, err := es.AppendActivity(a); err != nil {
t.Fatal(err)
}
if _, err := d.GetInterceptExecution(id); !errors.Is(err, ErrInterceptExecutionUnavailable) {
t.Fatal("cross-worker result was paired")
}
})
if _, err := d.GetInterceptExecution(approval(conv(), "", "missing", true)); !errors.Is(err, ErrInterceptSessionDeleted) {
t.Fatalf("missing call: %v", err)
}
for _, id := range []int64{approval(conv(), "", "", true), approval(conv(), "", "duplicate", false)} {
if _, err := d.GetInterceptExecution(id); !errors.Is(err, ErrInterceptExecutionUnavailable) {
t.Fatalf("invented execution for %d", id)
}
}
if got, err := d.GetInterceptExecution(-1); err != nil || got != nil {
t.Fatal("missing approval")
}
}
+141
View File
@@ -0,0 +1,141 @@
package db
import (
"slices"
"testing"
)
func TestInterceptApprovalFilters(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.Close() })
scope := "approval-filter-" + t.Name()
t.Cleanup(func() { _, _ = d.Exec(`DELETE FROM intercept_pending WHERE task_id IN ($1,$2)`, scope, scope+"-other") })
rule, err := d.CreateInterceptRule("filter fixture", "tool_name", "string", "Bash", "deny", "fixture", 1, true, false, 0, "deny")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _, _ = d.Exec(`DELETE FROM intercept_rules WHERE id=$1`, rule.ID) })
filter := InterceptApprovalFilter{Status: "denied", DecisionSource: "model"}
_, baseline, err := d.ListAllInterceptsPage(1, 20, filter)
if err != nil {
t.Fatal(err)
}
expected := []int64{}
for _, task := range []string{scope, scope + "-other"} {
copies := 3
if task != scope {
copies = 1
}
for _, status := range []string{"pending", "allowed", "denied", "timeout"} {
for _, source := range []string{"model", "rule", "unknown"} {
for i := 0; i < copies; i++ {
var ruleID int64
reason := "legacy reason"
if source == "rule" {
ruleID = rule.ID
}
if source == "model" {
reason = "[模型] fixture"
}
id, err := d.CreateDecidedIntercept(ruleID, 0, task, "test", "Bash", []byte(`{"command":"fixture"}`), status, reason)
if err != nil {
t.Fatal(err)
}
// Legacy source fields and tied timestamps must filter/page consistently.
if i == 0 {
if _, err = d.Exec(`UPDATE intercept_pending SET decision_source='' WHERE id=$1`, id); err != nil {
t.Fatal(err)
}
}
if task == scope && status == "denied" && source == "model" {
expected = append(expected, id)
}
}
}
}
}
if _, err := d.Exec(`UPDATE intercept_pending SET created_at='2099-01-01' WHERE task_id IN ($1,$2)`, scope, scope+"-other"); err != nil {
t.Fatal(err)
}
cases := []struct {
name string
filter InterceptApprovalFilter
want int
}{
{"all", InterceptApprovalFilter{}, 36},
}
for _, s := range []string{"pending", "allowed", "denied", "timeout"} {
cases = append(cases, struct {
name string
filter InterceptApprovalFilter
want int
}{s, InterceptApprovalFilter{Status: s}, 9})
}
for _, s := range []string{"model", "rule", "unknown"} {
cases = append(cases, struct {
name string
filter InterceptApprovalFilter
want int
}{s, InterceptApprovalFilter{DecisionSource: s}, 12})
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
rows, total, err := d.ListTaskInterceptsPage(scope, 1, 100, tc.filter)
if err != nil || total != tc.want || len(rows) != tc.want {
t.Fatalf("count=%d rows=%d err=%v", total, len(rows), err)
}
for _, r := range rows {
if r.TaskID == nil || *r.TaskID != scope || tc.filter.Status != "" && r.Status != tc.filter.Status || tc.filter.DecisionSource != "" && r.DecisionSource != tc.filter.DecisionSource {
t.Fatalf("nonmatching row: %+v", r)
}
}
})
}
slices.Reverse(expected)
var seen []int64
for page := 1; page <= 2; page++ {
rows, total, err := d.ListTaskInterceptsPage(scope, page, 2, filter)
if err != nil || total != 3 {
t.Fatalf("page %d total=%d err=%v", page, total, err)
}
for _, r := range rows {
seen = append(seen, r.ID)
}
}
if !slices.Equal(seen, expected) {
t.Fatalf("unstable combined-filter pagination: got %v want %v", seen, expected)
}
rows, total, err := d.ListTaskInterceptsPage(scope, 3, 2, filter)
if err != nil || len(rows) != 0 || total != 3 {
t.Fatalf("empty page count=%d rows=%d err=%v", total, len(rows), err)
}
_, total, err = d.ListAllInterceptsPage(1, 100, filter)
if err != nil || total != baseline+4 {
t.Fatalf("global filter count=%d want=%d err=%v", total, baseline+4, err)
}
// Values must remain SQL parameters even for direct store callers.
rows, total, err = d.ListTaskInterceptsPage(scope, 1, 20, InterceptApprovalFilter{Status: "denied' OR 1=1 --"})
if err != nil || total != 0 || len(rows) != 0 {
t.Fatalf("invalid value escaped filter: %d %v", total, err)
}
// A rule's persisted source survives deletion of its associated rule.
_, err = d.Exec(`UPDATE intercept_pending SET decision_source='rule' WHERE task_id=$1 AND rule_id=$2`, scope, rule.ID)
if err != nil {
t.Fatal(err)
}
if err = d.DeleteInterceptRule(rule.ID); err != nil {
t.Fatal(err)
}
rows, total, err = d.ListTaskInterceptsPage(scope, 1, 100, InterceptApprovalFilter{DecisionSource: "rule"})
if err != nil || total != 12 {
t.Fatalf("deleted-rule source count=%d err=%v", total, err)
}
for _, r := range rows {
if r.RuleID != nil || r.DecisionSource != "rule" {
t.Fatal("deleted rule lost its source")
}
}
}
+46
View File
@@ -0,0 +1,46 @@
package db
import (
"regexp"
"testing"
)
// 内置「删除类接口路径」规则匹配的是整个 tool_input JSON 串,因此用例直接以
// JSON 形态给出,与 Interceptor 实际拿到的 subject 一致。
func TestDeleteEndpointPathPattern(t *testing.T) {
re := regexp.MustCompile(deleteEndpointPathPattern)
hit := []string{
`{"command":"curl -s 'http://t.com/api/user/delete?id=1'"}`, // GET 打删除接口
`{"command":"curl -X POST http://t.com/admin/delete -d id=1"}`, // POST 打删除接口
`{"command":"curl 'http://t.com/api/deleteAll'"}`,
`{"command":"curl 'http://t.com/api/delete_user?id=1'"}`,
`{"command":"curl 'http://t.com/api/delete-user?id=1'"}`,
`{"url":"http://t.com/api/remove?id=1"}`,
`{"command":"curl http://t.com/files/unlink/3"}`,
`{"command":"curl http://t.com/api/del?id=2"}`,
`{"command":"curl -X POST http://t/v1/erase"}`,
`{"command":"curl http://t/admin/destroyAll"}`, // v1 的路径规则不允许后缀,这里补上
}
for _, s := range hit {
if !re.MatchString(s) {
t.Errorf("应命中却放行: %s", s)
}
}
// 动词后必须跟分隔符,避免 /delivery、/details 这类只读路径被误拦。
miss := []string{
`{"command":"curl 'http://t.com/api/delivery?id=1'"}`,
`{"command":"curl 'http://t.com/order/details'"}`,
`{"command":"curl 'http://t.com/api/delta/sync'"}`,
`{"command":"curl 'http://t.com/user/delegate'"}`,
`{"command":"curl 'http://delete.example.com/'"}`, // 删除动词出现在域名而非路径
`{"command":"curl 'http://t.com/remote/status'"}`,
`{"command":"nmap -p80 10.0.0.1"}`,
}
for _, s := range miss {
if re.MatchString(s) {
t.Errorf("误拦: %s", s)
}
}
}
+44
View File
@@ -0,0 +1,44 @@
package db
import (
"encoding/json"
"strings"
"testing"
)
func TestJsonbClean(t *testing.T) {
// A marshaled payload carrying a NUL byte (e.g. captured HTTP/tool output).
b, err := json.Marshal(map[string]string{"body": "ab\x00cd"})
if err != nil {
t.Fatal(err)
}
if !strings.Contains(string(b), `\u0000`) {
t.Fatalf("precondition: marshaled JSON should contain the NUL escape, got %s", b)
}
cleaned := jsonbClean(b)
if strings.Contains(string(cleaned), `\u0000`) {
t.Fatalf("jsonbClean left a NUL escape jsonb rejects: %s", cleaned)
}
// Result must stay valid JSON with the NUL simply dropped.
var out map[string]string
if err := json.Unmarshal(cleaned, &out); err != nil {
t.Fatalf("cleaned bytes are not valid JSON: %v (%s)", err, cleaned)
}
if out["body"] != "abcd" {
t.Fatalf("expected NUL stripped to \"abcd\", got %q", out["body"])
}
// A NUL escape typed literally in source text (doubled backslash) is preserved.
lit := []byte(`{"body":"\\u0000"}`)
if got := jsonbClean(lit); string(got) != string(lit) {
t.Fatalf("literal \\\\u0000 must be untouched, got %s", got)
}
// No NUL escape → returned unchanged.
plain := []byte(`{"body":"hello"}`)
if got := jsonbClean(plain); string(got) != string(plain) {
t.Fatalf("plain JSON must be untouched, got %s", got)
}
}
+62
View File
@@ -0,0 +1,62 @@
package db
import (
"testing"
)
// The raw-body columns were added after release, so existing installs only get
// them through llmRecordsMigrate. This runs the real migration against the dev
// PG on a table stripped back to its pre-upgrade shape, inside a transaction
// that is always rolled back — PG does transactional DDL, so nothing persists.
func TestLLMRecordsMigrateAddsRawColumnsToOldTable(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
if err := d.EnsureLLMRecordsTable(); err != nil {
t.Fatalf("ensure: %v", err)
}
tx, err := d.Begin()
if err != nil {
t.Fatalf("begin: %v", err)
}
defer tx.Rollback() //nolint:errcheck // the test never commits
// Roll the table back to how a pre-upgrade install looks.
if _, err := tx.Exec(`ALTER TABLE llm_records DROP COLUMN IF EXISTS raw_request, DROP COLUMN IF EXISTS raw_response`); err != nil {
t.Fatalf("simulate old table: %v", err)
}
if _, err := tx.Exec(llmRecordsMigrate); err != nil {
t.Fatalf("migrate: %v", err)
}
for _, col := range []string{"raw_request", "raw_response"} {
var n int
if err := tx.QueryRow(
`SELECT count(*) FROM information_schema.columns
WHERE table_name='llm_records' AND column_name=$1`, col).Scan(&n); err != nil {
t.Fatalf("inspect %s: %v", col, err)
}
if n != 1 {
t.Errorf("column %s missing after migrate", col)
}
}
// Re-running must stay a no-op (the migration runs on every startup).
if _, err := tx.Exec(llmRecordsMigrate); err != nil {
t.Fatalf("migrate is not idempotent: %v", err)
}
// An insert carrying raw bodies must round-trip through the migrated table.
var got string
if err := tx.QueryRow(
`INSERT INTO llm_records(model, raw_request, raw_response) VALUES ('m','{"a":1}','data: x')
RETURNING raw_response`).Scan(&got); err != nil {
t.Fatalf("insert into migrated table: %v", err)
}
if got != "data: x" {
t.Errorf("raw_response=%q want %q", got, "data: x")
}
}
+270
View File
@@ -0,0 +1,270 @@
package db
import (
"sort"
"strconv"
"time"
)
// LLMUsage is one lightweight LLM-call metering row — the always-on usage ledger,
// distinct from llm_records (which stores full request/response bodies and is a
// gated debug feature). One row per completion call, written on both success and
// error, so token accounting is complete even for interrupted/failed runs. Carries
// only the dimensions needed to slice token spend (model / profile / task / agent),
// never any prompt or response content.
type LLMUsage struct {
TaskID string `json:"task_id"` // task registry id (matches llm_records.task_id)
ExplorationID int64 `json:"exploration_id"` // exploration id parsed from the session (0 = unknown/non-task)
Worker string `json:"worker"` // agent lane: worker / planner / mainagent / goals
Model string `json:"model"`
ProfileName string `json:"profile_name"`
LatencyMs int `json:"latency_ms"`
InputTokens int `json:"input_tokens"`
OutputTokens int `json:"output_tokens"`
CacheRead int `json:"cache_read"`
CacheWrite int `json:"cache_write"`
Status string `json:"status"` // ok | error
}
const llmUsageSchema = `
CREATE TABLE IF NOT EXISTS llm_usage (
id BIGSERIAL PRIMARY KEY,
ts TIMESTAMPTZ NOT NULL DEFAULT now(),
task_id TEXT,
exploration_id BIGINT,
worker TEXT,
model TEXT,
profile_name TEXT,
latency_ms INTEGER,
input_tokens INTEGER NOT NULL DEFAULT 0,
output_tokens INTEGER NOT NULL DEFAULT 0,
cache_read INTEGER NOT NULL DEFAULT 0,
cache_write INTEGER NOT NULL DEFAULT 0,
status TEXT
);
CREATE INDEX IF NOT EXISTS idx_llm_usage_task ON llm_usage(task_id);
CREATE INDEX IF NOT EXISTS idx_llm_usage_model ON llm_usage(task_id, model);
CREATE INDEX IF NOT EXISTS idx_llm_usage_exp ON llm_usage(exploration_id);
`
// EnsureLLMUsageTable creates the llm_usage metering table if it does not exist.
func (d *DB) EnsureLLMUsageTable() error {
tx, err := d.Begin()
if err != nil {
return err
}
defer tx.Rollback() //nolint:errcheck
if err := coordinateWithSchemaMigration(tx); err != nil {
return err
}
if _, err := tx.Exec(llmUsageSchema); err != nil {
return err
}
return tx.Commit()
}
// InsertLLMUsage appends one metering row. Best-effort: callers log and continue on
// error (a lost metering row must never break the LLM call).
func (d *DB) InsertLLMUsage(u *LLMUsage) error {
var expID any
if u.ExplorationID > 0 {
expID = u.ExplorationID
}
_, err := d.Exec(`
INSERT INTO llm_usage(task_id, exploration_id, worker, model, profile_name, latency_ms, input_tokens, output_tokens, cache_read, cache_write, status)
VALUES (NULLIF($1,''),$2,NULLIF($3,''),NULLIF($4,''),NULLIF($5,''),$6,$7,$8,$9,$10,$11)`,
u.TaskID, expID, u.Worker, u.Model, u.ProfileName,
u.LatencyMs, u.InputTokens, u.OutputTokens, u.CacheRead, u.CacheWrite, u.Status)
return err
}
// TokenByModel aggregates a task's LLM token usage grouped by model, most-used
// first, from the always-on llm_usage ledger. taskID is the task registry id.
// Accurate even with per-agent model bindings, pool rotation/failover, and
// interrupted runs, since every call (success or error) is metered.
func (d *DB) TokenByModel(taskID string) ([]ModelTokenStat, error) {
rows, err := d.Query(`
SELECT COALESCE(NULLIF(model,''),'(unknown)') AS model, COUNT(*) AS calls,
COALESCE(SUM(input_tokens),0), COALESCE(SUM(output_tokens),0),
COALESCE(SUM(cache_read),0), COALESCE(SUM(cache_write),0)
FROM llm_usage
WHERE COALESCE(task_id,'') = $1
GROUP BY model
ORDER BY SUM(input_tokens) + SUM(output_tokens) DESC, model`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
out := []ModelTokenStat{}
for rows.Next() {
var m ModelTokenStat
if err := rows.Scan(&m.Model, &m.Calls, &m.InputTokens, &m.OutputTokens,
&m.CacheReadTokens, &m.CacheWriteTokens); err != nil {
return nil, err
}
out = append(out, m)
}
return out, rows.Err()
}
// ProfileUsage aggregates the whole ledger's token spend for one LLM profile
// (global, all tasks). Powers the dashboard's per-profile token card (new source).
type ProfileUsage struct {
ProfileName string `json:"profile_name"`
Calls int `json:"calls"`
Tasks int `json:"tasks"`
InputTokens int `json:"input_tokens"`
OutputTokens int `json:"output_tokens"`
CacheReadTokens int `json:"cache_read_tokens"`
CacheWriteTokens int `json:"cache_write_tokens"`
}
// UsageByProfile returns global token spend grouped by profile name, most-used
// first. profile_name may be empty for calls made on env/non-persisted configs.
func (d *DB) UsageByProfile() ([]ProfileUsage, error) {
rows, err := d.Query(`
SELECT COALESCE(profile_name,'') AS profile_name, COUNT(*) AS calls,
COUNT(DISTINCT task_id) AS tasks,
COALESCE(SUM(input_tokens),0), COALESCE(SUM(output_tokens),0),
COALESCE(SUM(cache_read),0), COALESCE(SUM(cache_write),0)
FROM llm_usage
GROUP BY profile_name
ORDER BY SUM(input_tokens) + SUM(output_tokens) DESC`)
if err != nil {
return nil, err
}
defer rows.Close()
out := []ProfileUsage{}
for rows.Next() {
var p ProfileUsage
if err := rows.Scan(&p.ProfileName, &p.Calls, &p.Tasks,
&p.InputTokens, &p.OutputTokens, &p.CacheReadTokens, &p.CacheWriteTokens); err != nil {
return nil, err
}
out = append(out, p)
}
if err := rows.Err(); err != nil {
return nil, err
}
archived, err := d.archivedTaskAggregates()
if err != nil {
return nil, err
}
byName := make(map[string]ProfileUsage, len(out))
for _, current := range out {
byName[current.ProfileName] = current
}
for _, aggregate := range archived {
for _, cold := range aggregate.TokenProfiles {
current := byName[cold.ProfileName]
current.ProfileName = cold.ProfileName
current.Calls += cold.Calls
current.Tasks += cold.Tasks
current.InputTokens += cold.InputTokens
current.OutputTokens += cold.OutputTokens
current.CacheReadTokens += cold.CacheReadTokens
current.CacheWriteTokens += cold.CacheWriteTokens
byName[cold.ProfileName] = current
}
}
out = out[:0]
for _, current := range byName {
out = append(out, current)
}
sort.Slice(out, func(i, j int) bool {
left := out[i].InputTokens + out[i].OutputTokens
right := out[j].InputTokens + out[j].OutputTokens
if left != right {
return left > right
}
return out[i].ProfileName < out[j].ProfileName
})
return out, nil
}
// ProfileDayUsage is one (profile, UTC calendar day) token bucket for the daily
// chart. Unlike the activity-based chart, ts is the real call time, so this is
// actual per-day consumption rather than tokens bucketed by task creation date.
type ProfileDayUsage struct {
ProfileName string `json:"profile_name"`
Date string `json:"date"` // YYYY-MM-DD (UTC)
InputTokens int `json:"input_tokens"`
OutputTokens int `json:"output_tokens"`
CacheReadTokens int `json:"cache_read_tokens"`
}
// UsageDaily returns per-(profile, day) token buckets for the past `days` days
// (default 365 when days<=0), so the dashboard can slice by profile + range.
func (d *DB) UsageDaily(days int) ([]ProfileDayUsage, error) {
if days <= 0 {
days = 365
}
rows, err := d.Query(`
SELECT COALESCE(profile_name,'') AS profile_name,
to_char(ts AT TIME ZONE 'UTC', 'YYYY-MM-DD') AS day,
COALESCE(SUM(input_tokens),0), COALESCE(SUM(output_tokens),0), COALESCE(SUM(cache_read),0)
FROM llm_usage
WHERE ts >= now() - ($1 * interval '1 day')
GROUP BY profile_name, day
ORDER BY day`, days)
if err != nil {
return nil, err
}
defer rows.Close()
out := []ProfileDayUsage{}
for rows.Next() {
var p ProfileDayUsage
if err := rows.Scan(&p.ProfileName, &p.Date, &p.InputTokens, &p.OutputTokens, &p.CacheReadTokens); err != nil {
return nil, err
}
out = append(out, p)
}
if err := rows.Err(); err != nil {
return nil, err
}
archived, err := d.archivedTaskAggregates()
if err != nil {
return nil, err
}
cutoff := time.Now().UTC().AddDate(0, 0, -days).Format("2006-01-02")
byKey := make(map[string]ProfileDayUsage, len(out))
for _, current := range out {
byKey[current.ProfileName+"\x00"+current.Date] = current
}
for _, aggregate := range archived {
for _, cold := range aggregate.TokenDaily {
if cold.Date < cutoff {
continue
}
key := cold.ProfileName + "\x00" + cold.Date
current := byKey[key]
current.ProfileName = cold.ProfileName
current.Date = cold.Date
current.InputTokens += cold.InputTokens
current.OutputTokens += cold.OutputTokens
current.CacheReadTokens += cold.CacheReadTokens
byKey[key] = current
}
}
out = out[:0]
for _, current := range byKey {
out = append(out, current)
}
sort.Slice(out, func(i, j int) bool {
if out[i].Date != out[j].Date {
return out[i].Date < out[j].Date
}
return out[i].ProfileName < out[j].ProfileName
})
return out, nil
}
// ParseExpID turns the exploration-id segment parsed from a session string into an
// int64 (0 when empty/non-numeric, e.g. chat sessions keyed by conversation id).
func ParseExpID(s string) int64 {
n, err := strconv.ParseInt(s, 10, 64)
if err != nil || n < 0 {
return 0
}
return n
}
+57
View File
@@ -0,0 +1,57 @@
package db
import "time"
// LLM failover health: one row per profile tracking its circuit-breaker state.
// The authoritative copy lives in memory (llmpool.Registry); this table only
// survives a restart so a cooling-off window isn't silently reset by one.
// LLMHealth is one profile's circuit-breaker state.
type LLMHealth struct {
ProfileID int64 `json:"profile_id"`
Fails int `json:"fails"` // consecutive failures; cleared on success
Trips int `json:"trips"` // total trips, drives the backoff ladder
OpenUntil *time.Time `json:"open_until"` // nil/past = closed (healthy)
LastError string `json:"last_error"`
LastAt time.Time `json:"last_at"`
}
// LoadLLMHealth returns the profiles still in an UNEXPIRED cooling-off window.
// Expired rows are deliberately skipped: after a restart a profile that has
// finished cooling should be treated as healthy again and re-probed on its next
// call, not resurrected as broken.
func (d *DB) LoadLLMHealth() ([]LLMHealth, error) {
rows, err := d.Query(`SELECT profile_id,fails,trips,open_until,COALESCE(last_error,''),last_at
FROM llm_profile_health WHERE open_until IS NOT NULL AND open_until > now()`)
if err != nil {
return nil, err
}
defer rows.Close()
var out []LLMHealth
for rows.Next() {
var h LLMHealth
if err := rows.Scan(&h.ProfileID, &h.Fails, &h.Trips, &h.OpenUntil, &h.LastError, &h.LastAt); err != nil {
return nil, err
}
out = append(out, h)
}
return out, rows.Err()
}
// SaveLLMHealth upserts one profile's circuit-breaker state.
func (d *DB) SaveLLMHealth(h LLMHealth) error {
_, err := d.Exec(`
INSERT INTO llm_profile_health(profile_id,fails,trips,open_until,last_error,last_at)
VALUES ($1,$2,$3,$4,$5,now())
ON CONFLICT (profile_id) DO UPDATE SET
fails=EXCLUDED.fails, trips=EXCLUDED.trips, open_until=EXCLUDED.open_until,
last_error=EXCLUDED.last_error, last_at=now()`,
h.ProfileID, h.Fails, h.Trips, h.OpenUntil, h.LastError)
return err
}
// ClearLLMHealth drops one profile's state (manual "recover now" from the UI).
func (d *DB) ClearLLMHealth(profileID int64) error {
_, err := d.Exec(`DELETE FROM llm_profile_health WHERE profile_id=$1`, profileID)
return err
}
+127
View File
@@ -0,0 +1,127 @@
package db
import (
"encoding/json"
"time"
)
// LLM 重试策略:五层重试的「次数 + 间隔」全局配置,见 docs/LLM重试设计.md。
// 存在 settings 表的一个 JSON 值里 —— 它是整机一份的运行参数,不值得为它开一张表;
// 读取走内置默认兜底,所以键不存在(全新库/从未配置过)时行为与写死常量时代完全一致。
const settingLLMRetryPolicy = "llm_retry_policy"
// RetryRule is one layer's knob pair. The zero value means "unset":
//
// Attempts 0 = 用内置默认次数; -1 = 关闭该层重试; >0 = 用该值
// IntervalMS 0 = 用该层原本的间隔策略(通常是指数退避); >0 = 改用固定毫秒间隔
//
// -1 是「显式关掉」而不是「0 次」,因为 0 已经被「未配置」占用了。
type RetryRule struct {
Attempts int `json:"attempts"`
IntervalMS int `json:"interval_ms"`
}
// Interval returns the configured fixed interval, or 0 when unset (caller keeps
// its own default ladder).
func (r RetryRule) Interval() time.Duration {
if r.IntervalMS <= 0 {
return 0
}
return time.Duration(r.IntervalMS) * time.Millisecond
}
// Or returns the rule with each unset field filled in from fallback. Used to
// layer a profile override on top of the global policy field by field, so a
// profile that only pins the interval still inherits the global count.
func (r RetryRule) Or(fallback RetryRule) RetryRule {
if r.Attempts == 0 {
r.Attempts = fallback.Attempts
}
if r.IntervalMS == 0 {
r.IntervalMS = fallback.IntervalMS
}
return r
}
// retry knob bounds. A count above the cap turns a blip into a token bonfire;
// an interval above an hour outlives any transient failure worth waiting out.
const (
maxRetryAttempts = 20
maxRetryIntervalMS = 3600_000 // 1h
)
// Clamped returns the rule with out-of-range values pulled back into the sane
// band (attempts within [-1, 20], interval within [0, 1h]).
func (r RetryRule) Clamped() RetryRule {
if r.Attempts < -1 {
r.Attempts = -1
}
if r.Attempts > maxRetryAttempts {
r.Attempts = maxRetryAttempts
}
if r.IntervalMS < 0 {
r.IntervalMS = 0
}
if r.IntervalMS > maxRetryIntervalMS {
r.IntervalMS = maxRetryIntervalMS
}
return r
}
// Clamped bounds a profile's override the same way the global policy is bounded,
// so a hand-crafted API payload can't land a value the CHECK constraint rejects.
func (o RetryOverride) Clamped() RetryOverride {
o.Connect, o.Empty, o.Stream = o.Connect.Clamped(), o.Empty.Clamped(), o.Stream.Clamped()
return o
}
// LLMRetryPolicy holds the五层 retry configuration. Connect/Empty/Stream are the
// per-request layers (a profile may override them, see LLMProfile.Retry);
// Breaker and Intent are process-wide by nature and live only here.
type LLMRetryPolicy struct {
// Connect:SDK 建连重试(连接重置/超时/429/5xx,流开始前)。默认 3 次、指数退避。
Connect RetryRule `json:"connect"`
// Empty:SDK 空响应重试(完成但无 content block,仅 openai 格式)。默认 2 次、指数退避。
Empty RetryRule `json:"empty"`
// Stream:同 provider 安全窗口重试(未交付输出前的断流重放)。默认 2 次、0.5s 起指数(封顶 4s)。
Stream RetryRule `json:"stream"`
// Breaker:轮询熔断。Attempts=连续几次瞬时失败触发熔断(默认 3,-1=瞬时失败不熔断,
// 硬失败如余额不足/密钥失效仍立即熔断);IntervalMS=固定冷却时长(0=默认 1/5/30min 梯度)。
Breaker RetryRule `json:"breaker"`
// Intent:worker 以 model_error 收场后的整条意图重跑。默认 2 次、固定 3s。
Intent RetryRule `json:"intent"`
}
// Clamped returns the policy with every rule clamped.
func (p LLMRetryPolicy) Clamped() LLMRetryPolicy {
p.Connect, p.Empty, p.Stream = p.Connect.Clamped(), p.Empty.Clamped(), p.Stream.Clamped()
p.Breaker, p.Intent = p.Breaker.Clamped(), p.Intent.Clamped()
return p
}
// LLMRetryPolicy reads the global retry policy. A missing or unparseable value
// yields the zero policy — i.e. every layer on its built-in default.
func (d *DB) LLMRetryPolicy() LLMRetryPolicy {
var p LLMRetryPolicy
if d == nil {
return p
}
raw, ok, err := d.GetSetting(settingLLMRetryPolicy)
if err != nil || !ok || raw == "" {
return p
}
if err := json.Unmarshal([]byte(raw), &p); err != nil {
return LLMRetryPolicy{}
}
return p.Clamped()
}
// SetLLMRetryPolicy persists the global retry policy (values are clamped first).
func (d *DB) SetLLMRetryPolicy(p LLMRetryPolicy) error {
raw, err := json.Marshal(p.Clamped())
if err != nil {
return err
}
return d.SetSetting(settingLLMRetryPolicy, string(raw))
}
+145
View File
@@ -0,0 +1,145 @@
package db
import (
"testing"
"time"
)
// A profile's retry override must survive a full round trip through the real
// column list — this is what catches a mis-ordered scan/insert after adding six
// columns at once. Also pins the "empty api key on update keeps the row usable"
// path, since that UPDATE has its own parameter numbering.
func TestProfileRetryRoundTrip(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
want := RetryOverride{
Connect: RetryRule{Attempts: 5, IntervalMS: 2000},
Empty: RetryRule{Attempts: -1},
Stream: RetryRule{IntervalMS: 250},
}
id, err := d.SaveProfile(&LLMProfile{
Name: "t-retry-roundtrip", Format: "openai", Model: "m", APIKey: "k", Retry: want,
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { d.Exec(`DELETE FROM llm_profiles WHERE id=$1`, id) })
got, err := d.ProfileByID(id)
if err != nil || got == nil {
t.Fatalf("ProfileByID: %v", err)
}
if got.Retry != want {
t.Fatalf("retry=%+v, want %+v", got.Retry, want)
}
// Keyless update path (the UI sends no key when the user didn't retype it).
got.APIKey = ""
got.Retry.Stream = RetryRule{Attempts: 3, IntervalMS: 700}
if _, err := d.SaveProfile(got); err != nil {
t.Fatal(err)
}
after, err := d.ProfileByID(id)
if err != nil || after == nil {
t.Fatalf("ProfileByID after update: %v", err)
}
if after.Retry.Stream != (RetryRule{Attempts: 3, IntervalMS: 700}) {
t.Fatalf("stream=%+v after update", after.Retry.Stream)
}
if after.Retry.Connect != want.Connect || after.Retry.Empty != want.Empty {
t.Fatalf("untouched rules changed: %+v", after.Retry)
}
// The listing query reads a different column list — it must agree.
profiles, err := d.ListProfiles()
if err != nil {
t.Fatal(err)
}
for _, p := range profiles {
if p.ID == id && p.Retry.Connect != want.Connect {
t.Fatalf("ListProfiles retry=%+v, want %+v", p.Retry, want)
}
}
}
// Out-of-range values are clamped on the way in, so the DB CHECK constraint is
// never what the user hears about.
func TestProfileRetryClampedOnSave(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
id, err := d.SaveProfile(&LLMProfile{
Name: "t-retry-clamp", Format: "openai", Model: "m", APIKey: "k",
Retry: RetryOverride{Connect: RetryRule{Attempts: -99, IntervalMS: -5}},
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { d.Exec(`DELETE FROM llm_profiles WHERE id=$1`, id) })
got, err := d.ProfileByID(id)
if err != nil || got == nil {
t.Fatalf("ProfileByID: %v", err)
}
if got.Retry.Connect != (RetryRule{Attempts: -1}) {
t.Fatalf("connect=%+v, want attempts -1 and no interval", got.Retry.Connect)
}
}
// The global policy round-trips through settings, and an unset key reads back as
// "everything on its built-in default".
func TestLLMRetryPolicyRoundTrip(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
before, hadBefore, err := d.GetSetting(settingLLMRetryPolicy)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
if hadBefore {
d.SetSetting(settingLLMRetryPolicy, before)
} else {
d.Exec(`DELETE FROM settings WHERE key=$1`, settingLLMRetryPolicy)
}
})
want := LLMRetryPolicy{
Connect: RetryRule{Attempts: 4, IntervalMS: 1500},
Breaker: RetryRule{Attempts: 2, IntervalMS: 90_000},
Intent: RetryRule{Attempts: -1},
}
if err := d.SetLLMRetryPolicy(want); err != nil {
t.Fatal(err)
}
got := d.LLMRetryPolicy()
if got != want {
t.Fatalf("policy=%+v, want %+v", got, want)
}
if d := got.Breaker.Interval(); d != 90*time.Second {
t.Fatalf("breaker interval=%v, want 90s", d)
}
// 越界值写进去也会被夹回区间,读出来是夹紧后的值。
if err := d.SetLLMRetryPolicy(LLMRetryPolicy{Stream: RetryRule{Attempts: 999, IntervalMS: 99_999_999}}); err != nil {
t.Fatal(err)
}
if got := d.LLMRetryPolicy().Stream; got.Attempts != 20 || got.IntervalMS != 3600_000 {
t.Fatalf("stream=%+v, want the 20 / 1h caps", got)
}
// 键不存在 = 全默认。
if _, err := d.Exec(`DELETE FROM settings WHERE key=$1`, settingLLMRetryPolicy); err != nil {
t.Fatal(err)
}
if got := d.LLMRetryPolicy(); got != (LLMRetryPolicy{}) {
t.Fatalf("unset policy=%+v, want zero", got)
}
}
+78
View File
@@ -0,0 +1,78 @@
package db
import (
"context"
"time"
)
// DBLog is one persisted backend log row from the server_logs table.
type DBLog struct {
ID int64
CreatedAt time.Time
Level string
Tag string
Text string
}
// InsertLog appends one log line and returns its auto-assigned id.
func (d *DB) InsertLog(level, tag, text string) (int64, error) {
var id int64
err := d.QueryRowContext(context.Background(),
"INSERT INTO server_logs(level,tag,text) VALUES($1,$2,$3) RETURNING id",
level, tag, text,
).Scan(&id)
return id, err
}
// RecentLogs returns the most recent limit rows, oldest-first.
func (d *DB) RecentLogs(limit int) ([]*DBLog, error) {
rows, err := d.QueryContext(context.Background(),
`SELECT id, created_at, level, tag, text
FROM server_logs
ORDER BY id DESC
LIMIT $1`, limit)
if err != nil {
return nil, err
}
defer rows.Close()
var out []*DBLog
for rows.Next() {
l := &DBLog{}
if err := rows.Scan(&l.ID, &l.CreatedAt, &l.Level, &l.Tag, &l.Text); err != nil {
return nil, err
}
out = append(out, l)
}
// Reverse to oldest-first
for i, j := 0, len(out)-1; i < j; i, j = i+1, j-1 {
out[i], out[j] = out[j], out[i]
}
return out, nil
}
// ListLogsBefore returns up to limit rows with id < beforeID, oldest-first.
func (d *DB) ListLogsBefore(beforeID int64, limit int) ([]*DBLog, error) {
rows, err := d.QueryContext(context.Background(),
`SELECT id, created_at, level, tag, text
FROM server_logs
WHERE id < $1
ORDER BY id DESC
LIMIT $2`, beforeID, limit)
if err != nil {
return nil, err
}
defer rows.Close()
var out []*DBLog
for rows.Next() {
l := &DBLog{}
if err := rows.Scan(&l.ID, &l.CreatedAt, &l.Level, &l.Tag, &l.Text); err != nil {
return nil, err
}
out = append(out, l)
}
// Reverse to oldest-first
for i, j := 0, len(out)-1; i < j; i, j = i+1, j-1 {
out[i], out[j] = out[j], out[i]
}
return out, nil
}
+134
View File
@@ -0,0 +1,134 @@
package db
import (
"net"
"net/url"
"regexp"
"strconv"
"strings"
"golang.org/x/net/publicsuffix"
)
// 归一化自然键 (nkey):移植自旧 graph/id.go,去掉 StableID 哈希(PG 用 BIGSERIAL 主键 +
// UNIQUE(type, nkey) 去重)。子资产的 nkey 内嵌父资产的 int64 id,把层级编码进键。
func DomainKey(fqdn string) string {
return strings.TrimSuffix(strings.ToLower(strings.TrimSpace(fqdn)), ".")
}
func IPKey(ip string) string { return strings.TrimSpace(ip) }
// RootDomain returns the registrable domain (eTLD+1) for a host and whether the
// host itself IS that apex (§3.1). Edge cases (§3.1 边界处理): an IP literal or a
// host publicsuffix can't classify (localhost / internal / non-ICANN TLD) is
// returned unchanged as its own root with isApex=true — best-effort, never treated
// as a subdomain.
func RootDomain(host string) (root string, isApex bool) {
h := DomainKey(host)
if h == "" || net.ParseIP(h) != nil {
return h, true
}
etld1, err := publicsuffix.EffectiveTLDPlusOne(h)
if err != nil || etld1 == "" {
return h, true
}
return etld1, h == etld1
}
func PortKey(ipID int64, proto string, port int) string {
return itoa(ipID) + "|" + strings.ToLower(proto) + "|" + strconv.Itoa(port)
}
func ServiceKey(portID int64, svcName string) string {
return itoa(portID) + "|" + strings.ToLower(svcName)
}
func SiteKey(scheme, host string, port int) string {
return strings.ToLower(scheme) + "|" + strings.ToLower(host) + "|" + strconv.Itoa(port)
}
func EndpointKey(siteID int64, method, urlTemplate string) string {
return itoa(siteID) + "|" + strings.ToUpper(method) + "|" + urlTemplate
}
func ParameterKey(endpointID int64, location, name string) string {
return itoa(endpointID) + "|" + strings.ToLower(location) + "|" + name
}
// NormalizeParamName 归一化参数名(endpoint.params 元素的「相同引用」判定)。
// 规则:lower + trim,不做同义词合并(userId/user_id/uid 视为不同)。写入与查询共享此实现,
// 保证「按参数名查同公司接口」可复现。
func NormalizeParamName(name string) string {
return strings.ToLower(strings.TrimSpace(name))
}
func TechKey(name, version string) string {
return strings.ToLower(name) + "|" + version
}
func itoa(n int64) string { return strconv.FormatInt(n, 10) }
var (
reNumeric = regexp.MustCompile(`^\d+$`)
reUUID = regexp.MustCompile(`^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$`)
reHex = regexp.MustCompile(`^[0-9a-fA-F]{16,}$`)
reLong = regexp.MustCompile(`^[A-Za-z0-9_-]{24,}$`)
)
// TemplatePath collapses high-cardinality path segments into placeholders so the
// graph is not flooded by instances: /user/123 -> /user/{id}.
func TemplatePath(path string) string {
if path == "" {
return "/"
}
segs := strings.Split(path, "/")
for i, s := range segs {
switch {
case s == "":
continue
case reNumeric.MatchString(s):
segs[i] = "{id}"
case reUUID.MatchString(s):
segs[i] = "{uuid}"
case reHex.MatchString(s):
segs[i] = "{hex}"
case reLong.MatchString(s):
segs[i] = "{token}"
}
}
return strings.Join(segs, "/")
}
// SplitURL parses a raw URL into scheme/host/port/urlTemplate/params for building
// site/endpoint/param keys.
func SplitURL(raw, method string) (scheme, host string, port int, urlTemplate string, params []string, err error) {
u, err := url.Parse(raw)
if err != nil {
return "", "", 0, "", nil, err
}
scheme = strings.ToLower(u.Scheme)
host = strings.ToLower(u.Hostname())
port = defaultPort(scheme, u.Port())
urlTemplate = TemplatePath(u.EscapedPath())
for k := range u.Query() {
params = append(params, k)
}
return scheme, host, port, urlTemplate, params, nil
}
func defaultPort(scheme, p string) int {
if p != "" {
if n, err := strconv.Atoi(p); err == nil {
return n
}
}
switch scheme {
case "https":
return 443
case "http":
return 80
default:
return 0
}
}
+551
View File
@@ -0,0 +1,551 @@
package db
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"log"
"strings"
"time"
"github.com/Autumn-27/artex/notify"
)
// 本文件是 IM 推送的渠道配置与事件层。投递任务的领取与状态流转见
// db/notification_delivery.go。
//
// 两条不变量,改这个文件时务必保持:
//
// 1. 写漏洞的事务(RecordFindingTx)只调用 InsertNotificationEventTx 做一次盲插,
// 不读任何通知相关的表、不做过滤匹配。任何在这里引入的读操作都可能因为
// 用户配错的过滤条件而污染甚至中止漏洞写入事务。
// 2. 过滤匹配永不报错:配置畸形一律按「命中」处理(见 notify.Match)。宁可多推,
// 不可漏推。
// ErrNotificationChannelNotFound 渠道不存在。
var ErrNotificationChannelNotFound = errors.New("通知渠道不存在")
// 投递状态。
const (
NotifyStatePending = "pending" // 待发
NotifyStateSending = "sending" // 已被某个 dispatcher 领取,租约未到期
NotifyStateSent = "sent" // 已送达
NotifyStateFailed = "failed" // 重试耗尽或永久失败,可手动重发
NotifyStateSkipped = "skipped" // 渠道已停用,不再发送
)
// 推送模式。
const (
NotifyModeRealtime = "realtime"
NotifyModeDigest = "digest"
)
// ValidNotifyMode 白名单校验推送模式(与 findings.status 同理:不用 DB CHECK,
// 便于后续扩展)。
func ValidNotifyMode(m string) bool {
return m == NotifyModeRealtime || m == NotifyModeDigest
}
// NotificationChannel 是一个渠道实例配置。Config 与 Filter 保持原始 JSON,
// 解析交给 notify 包——db 层不理解它们的字段含义。
type NotificationChannel struct {
ID int64 `json:"id"`
Name string `json:"name"`
Kind string `json:"kind"`
Mode string `json:"mode"`
Config json.RawMessage `json:"config"`
Filter json.RawMessage `json:"filter"`
// Enabled 用指针是为了区分「没传这个字段」与「显式传 false」——
// 前端开关控件只提交被改动的字段。
Enabled *bool `json:"enabled,omitempty"`
RatePerMin int `json:"rate_per_min"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// IsEnabled 返回渠道是否启用;Enabled 为 nil(未加载)时按启用处理。
func (c *NotificationChannel) IsEnabled() bool { return c.Enabled == nil || *c.Enabled }
// NotificationEvent 是一条事件事实。
type NotificationEvent struct {
ID int64 `json:"id"`
Kind string `json:"kind"`
FindingID int64 `json:"finding_id"`
Snapshot json.RawMessage `json:"snapshot"`
CreatedAt time.Time `json:"created_at"`
}
const notificationChannelCols = `id, name, kind, enabled, config, mode, filter, rate_per_min, created_at, updated_at`
func scanNotificationChannel(sc interface{ Scan(...any) error }) (*NotificationChannel, error) {
var c NotificationChannel
var enabled bool
if err := sc.Scan(&c.ID, &c.Name, &c.Kind, &enabled, &c.Config, &c.Mode, &c.Filter, &c.RatePerMin, &c.CreatedAt, &c.UpdatedAt); err != nil {
return nil, err
}
c.Enabled = &enabled
return &c, nil
}
// ListNotificationChannels 返回全部渠道实例,启用的排在前面、同级按 id。
// 排序放在 SQL 里是为了让 UI 与 dispatcher 看到同一个稳定顺序。
func (d *DB) ListNotificationChannels(ctx context.Context) ([]*NotificationChannel, error) {
rows, err := d.QueryContext(ctx, `SELECT `+notificationChannelCols+` FROM notification_channels
ORDER BY enabled DESC, id`)
if err != nil {
return nil, err
}
defer rows.Close()
out := []*NotificationChannel{}
for rows.Next() {
c, err := scanNotificationChannel(rows)
if err != nil {
return nil, err
}
out = append(out, c)
}
return out, rows.Err()
}
// NotificationChannelByID 取单个渠道。
func (d *DB) NotificationChannelByID(ctx context.Context, id int64) (*NotificationChannel, error) {
row := d.QueryRowContext(ctx, `SELECT `+notificationChannelCols+` FROM notification_channels WHERE id=$1`, id)
c, err := scanNotificationChannel(row)
if err == sql.ErrNoRows {
return nil, ErrNotificationChannelNotFound
}
return c, err
}
// SaveNotificationChannel 新建或更新一个渠道。
//
// 更新时只覆盖调用方显式给出的字段(非 nil / 非空),这样前端可以提交局部
// 修改的抽屉表单,而不必回传 config 里那些它没展示的字段——回传反而会造成
// 「掩码值把真实密钥覆盖掉」的事故。
func (d *DB) SaveNotificationChannel(ctx context.Context, c *NotificationChannel) (int64, error) {
if c.Mode == "" {
c.Mode = NotifyModeRealtime
}
// 这里刻意**不**对 0 做任何加工:0 是合法配置,含义是「不限流」。
//
// 曾经写成 `if c.RatePerMin <= 0 { c.RatePerMin = 默认值 }`,本意是「未指定时
// 给个安全默认」,但那把「显式设成 0」也一起吞掉了——文档、UI 提示与
// takeTokens 都把 0 解释为不限流,唯独这里悄悄改成 20(钉钉/企微/Telegram)
// 或 100(飞书),操作者以为放开了限流、实际被 20/分钟卡着且没有任何提示。
//
// 「未指定」与「显式 0」的区别只有调用方知道(请求体里字段缺省 vs 明确传 0),
// 所以默认值由 server 层在字段缺省时填,见 notifyCreateChannel。
if c.RatePerMin < 0 {
return 0, errors.New("限流值不能为负")
}
if c.Config == nil {
c.Config = json.RawMessage(`{}`)
}
if c.Filter == nil {
c.Filter = json.RawMessage(`{}`)
}
enabled := c.IsEnabled()
if c.ID == 0 {
var id int64
err := d.QueryRowContext(ctx, `INSERT INTO notification_channels(name,kind,enabled,config,mode,filter,rate_per_min)
VALUES($1,$2,$3,$4,$5,$6,$7) RETURNING id`,
c.Name, c.Kind, enabled, string(c.Config), c.Mode, string(c.Filter), c.RatePerMin).Scan(&id)
return id, err
}
res, err := d.ExecContext(ctx, `UPDATE notification_channels
SET name=$2, kind=$3, enabled=$4, config=$5, mode=$6, filter=$7, rate_per_min=$8
WHERE id=$1`,
c.ID, c.Name, c.Kind, enabled, string(c.Config), c.Mode, string(c.Filter), c.RatePerMin)
if err != nil {
return 0, err
}
if n, _ := res.RowsAffected(); n == 0 {
return 0, ErrNotificationChannelNotFound
}
return c.ID, nil
}
// SetNotificationChannelEnabled 切换启停。
//
// 停用一个渠道时,把它尚未发出的投递一并标记为 skipped:否则重新启用后
// 会突然收到一批「停用期间积压」的旧漏洞,时效已失且容易误判为新增。
func (d *DB) SetNotificationChannelEnabled(ctx context.Context, id int64, enabled bool) error {
return d.WithEvidenceTx(ctx, func(tx *sql.Tx) error {
res, err := tx.ExecContext(ctx, `UPDATE notification_channels SET enabled=$2 WHERE id=$1`, id, enabled)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotificationChannelNotFound
}
if !enabled {
if _, err := tx.ExecContext(ctx, `UPDATE notification_deliveries SET state=$2, last_error=$3
WHERE channel_id=$1 AND state IN ($4,$5)`,
id, NotifyStateSkipped, "渠道已停用", NotifyStatePending, NotifyStateSending); err != nil {
return err
}
}
return nil
})
}
// DeleteNotificationChannel 删除渠道。其投递历史随外键级联删除
// (渠道配置都没了,历史无从解读)。
func (d *DB) DeleteNotificationChannel(ctx context.Context, id int64) error {
res, err := d.ExecContext(ctx, `DELETE FROM notification_channels WHERE id=$1`, id)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotificationChannelNotFound
}
return nil
}
// RecordNotificationEventTx 在调用方的事务里**尽力**写入一条推送事件。
//
// 这是漏洞写入路径上唯一的通知相关改动:一次 INSERT,不读任何表、不认识渠道、
// 不跑过滤。事务提交即保证「漏洞落库」与「推送任务存在」原子一致,
// 不存在提交成功却没入队、消息永久丢失的窗口。
//
// 两个关键设计,都不是随手写的:
//
// 1. **为什么用 SAVEPOINT**:PostgreSQL 里事务内任一语句报错会让整个事务进入
// aborted 状态,此后所有语句(含 COMMIT)一律失败。所以「忽略这条 INSERT
// 的错误、让调用方继续提交」在 PG 里是做不到的——除非用保存点把错误隔离在
// 这一条语句上。没有保存点,就只剩「整笔回滚」这一个选项。
//
// 2. **为什么整笔回滚是错的**:推送是便利功能,漏洞记录才是产品本身。一个通知
// 表的问题(旧库未迁移、磁盘瞬时故障)不该让高危漏洞存不进库。所以这里隔离
// 错误、记日志、返回 false,让漏洞写入照常提交——代价是丢掉这一条推送。
// 返回 bool 而非 error 是刻意的:调用方不该把它当作会影响写入成败的错误。
func RecordNotificationEventTx(ctx context.Context, tx *sql.Tx, kind string, findingID int64, snap notify.Snapshot) bool {
raw, err := json.Marshal(snap)
if err != nil {
log.Printf("[notify] 序列化推送事件失败 finding=%d: %v", findingID, err)
return false
}
if _, err := tx.ExecContext(ctx, `SAVEPOINT notify_event`); err != nil {
log.Printf("[notify] 建立保存点失败 finding=%d: %v", findingID, err)
return false
}
if _, err := tx.ExecContext(ctx, `INSERT INTO notification_events(kind,finding_id,snapshot) VALUES($1,$2,$3)`,
kind, findingID, string(raw)); err != nil {
log.Printf("[notify] 写入推送事件失败 finding=%d(漏洞记录不受影响): %v", findingID, err)
// 回滚到保存点,把事务从 aborted 状态里救回来。
if _, rbErr := tx.ExecContext(ctx, `ROLLBACK TO SAVEPOINT notify_event`); rbErr != nil {
log.Printf("[notify] 回滚到保存点失败 finding=%d: %v", findingID, rbErr)
}
return false
}
// 释放保存点,避免长事务里积攒无用的保存点。
_, _ = tx.ExecContext(ctx, `RELEASE SAVEPOINT notify_event`)
return true
}
// AddNotificationEvent 是 InsertNotificationEventTx 的独立事务版本,供不在
// 既有事务里的调用点使用(如渠道的「发送测试消息」,它没有真实 finding)。
func (d *DB) AddNotificationEvent(ctx context.Context, kind string, findingID int64, snap notify.Snapshot) (int64, error) {
raw, err := json.Marshal(snap)
if err != nil {
return 0, fmt.Errorf("序列化通知事件快照失败: %w", err)
}
var id int64
err = d.QueryRowContext(ctx, `INSERT INTO notification_events(kind,finding_id,snapshot) VALUES($1,$2,$3) RETURNING id`,
kind, findingID, string(raw)).Scan(&id)
return id, err
}
// FanOutPendingEvents 把尚未分派的漏洞事件按当前启用的渠道展开成投递任务,
// 返回本轮处理的事件数与新建的投递数。
//
// 整轮操作在一个事务里:事件用 FOR UPDATE SKIP LOCKED 领取,多个进程同时跑
// 也各自领到不同的行(项目里归档队列的领取用的是同一套手法,见
// db/task_archives.go 的 completeNextArchiveJob)。
//
// 过滤匹配刻意放在 Go 侧而非 SQL:渠道的过滤条件是一组可选字段的 JSONB,
// 用 SQL 表达六种组合的匹配会让查询难以维护,而渠道数量是「人手配的几条」,
// 全量加载后在内存里逐条比对更快也更好测。
//
// 未命中任何渠道的事件同样会被标记 fanned_out ——否则它会永远留在待分派集合里,
// 每个 tick 被重扫一遍。
func (d *DB) FanOutPendingEvents(ctx context.Context, limit int) (eventCount, deliveryCount int, err error) {
if limit <= 0 {
limit = 200
}
tx, err := d.BeginTx(ctx, nil)
if err != nil {
return 0, 0, err
}
defer tx.Rollback() //nolint:errcheck // 提交成功后是 no-op
channels, err := listEnabledNotificationChannelsTx(ctx, tx)
if err != nil {
return 0, 0, err
}
rows, err := tx.QueryContext(ctx, `SELECT id, kind, finding_id, snapshot FROM notification_events
WHERE NOT fanned_out ORDER BY id FOR UPDATE SKIP LOCKED LIMIT $1`, limit)
if err != nil {
return 0, 0, err
}
var (
events []NotificationEvent
parsedSnaps []notify.Snapshot
)
for rows.Next() {
var ev NotificationEvent
if err := rows.Scan(&ev.ID, &ev.Kind, &ev.FindingID, &ev.Snapshot); err != nil {
rows.Close()
return 0, 0, err
}
var snap notify.Snapshot
// 快照是我们自己写的,理论上必定可解析;解析失败不阻断投递流程,
// 但这条事件会因字段全空而被所有带过滤条件的渠道跳过——宁可少推一条
// 也不让一个坏行卡死整个队列。
_ = json.Unmarshal(ev.Snapshot, &snap)
// kind 以行内值为准:快照里那份是渲染用的副本,可能被旧版本写过。
snap.Kind = ev.Kind
events = append(events, ev)
parsedSnaps = append(parsedSnaps, snap)
}
rows.Close()
if err := rows.Err(); err != nil {
return 0, 0, err
}
if len(events) == 0 {
return 0, 0, tx.Commit()
}
type pending struct {
eventID int64
channelID int64
}
var toInsert []pending
for i, snap := range parsedSnaps {
for _, ch := range channels {
if !notify.Match(notify.ParseFilter(ch.Filter), snap) {
continue
}
toInsert = append(toInsert, pending{eventID: events[i].ID, channelID: ch.ID})
}
}
if len(toInsert) > 0 {
var (
vals []string
args []any
)
for _, p := range toInsert {
vals = append(vals, fmt.Sprintf("($%d,$%d)", len(args)+1, len(args)+2))
args = append(args, p.eventID, p.channelID)
}
if _, err := tx.ExecContext(ctx, `INSERT INTO notification_deliveries(event_id,channel_id) VALUES `+strings.Join(vals, ","), args...); err != nil {
return 0, 0, err
}
}
// 标记本轮事件已分派。未命中任何渠道的事件也一起标记(见函数注释)。
ids := make([]string, 0, len(events))
markArgs := make([]any, 0, len(events))
for _, ev := range events {
markArgs = append(markArgs, ev.ID)
ids = append(ids, fmt.Sprintf("$%d", len(markArgs)))
}
if _, err := tx.ExecContext(ctx, `UPDATE notification_events SET fanned_out=true WHERE id IN (`+strings.Join(ids, ",")+`)`, markArgs...); err != nil {
return 0, 0, err
}
return len(events), len(toInsert), tx.Commit()
}
// listEnabledNotificationChannelsTx 在事务里取启用中的渠道。数量很少,
// 不做分页也不加缓存——缓存会引入「改了配置何时生效」这个额外的时序问题。
func listEnabledNotificationChannelsTx(ctx context.Context, tx *sql.Tx) ([]*NotificationChannel, error) {
rows, err := tx.QueryContext(ctx, `SELECT id, name, kind, config, mode, filter, rate_per_min
FROM notification_channels WHERE enabled ORDER BY id`)
if err != nil {
return nil, err
}
defer rows.Close()
out := []*NotificationChannel{}
for rows.Next() {
var c NotificationChannel
if err := rows.Scan(&c.ID, &c.Name, &c.Kind, &c.Config, &c.Mode, &c.Filter, &c.RatePerMin); err != nil {
return nil, err
}
out = append(out, &c)
}
return out, rows.Err()
}
// NotificationAssetNames 把资产 id 解析成简短展示名,供推送消息使用。
//
// 返回顺序与入参一致、长度可能小于入参(不存在的 id 被跳过)。保持入参顺序是
// 为了让同一条漏洞的消息在多次投递里资产顺序稳定——否则重试后收到的消息里
// 资产次序变了,会被误读成「资产变了」。
func (d *DB) NotificationAssetNames(ctx context.Context, ids []int64) ([]string, error) {
if len(ids) == 0 {
return nil, nil
}
ph, args := placeholders(1, ids)
rows, err := d.QueryContext(ctx, `SELECT id, type, domain, ip, url, app_name, bundle_id FROM assets WHERE id IN (`+ph+`)`, args...)
if err != nil {
return nil, err
}
defer rows.Close()
labels := map[int64]string{}
for rows.Next() {
var (
id int64
typ string
domain, ip, url sql.NullString
appName, bundleID sql.NullString
)
if err := rows.Scan(&id, &typ, &domain, &ip, &url, &appName, &bundleID); err != nil {
return nil, err
}
labels[id] = assetDisplayName(typ, domain.String, ip.String, url.String, appName.String, bundleID.String)
}
if err := rows.Err(); err != nil {
return nil, err
}
out := make([]string, 0, len(ids))
seen := map[int64]bool{}
for _, id := range ids {
if seen[id] {
continue
}
seen[id] = true
if label, ok := labels[id]; ok && label != "" {
out = append(out, label)
}
}
return out, nil
}
// assetDisplayName 按资产类型挑选最具辨识度的标识。
// 兜底返回空串,由调用方决定怎么呈现「名字取不到的资产」——本函数不臆造占位符,
// 否则「资产#42」这种噪音会混进推送消息里,读者还以为是真实域名。
func assetDisplayName(typ, domain, ip, url, appName, bundleID string) string {
pick := func(vals ...string) string {
for _, v := range vals {
if strings.TrimSpace(v) != "" {
return v
}
}
return ""
}
switch typ {
case "root_domain", "subdomain":
return domain
case "ip":
return ip
case "app":
return pick(appName, bundleID)
case "service", "endpoint":
return pick(url, domain, ip)
default:
return pick(domain, ip, url, appName)
}
}
// SetFindingStatusWithNotify 更新漏洞处置状态,并在同一事务里登记一条状态变更
// 推送事件。
//
// 返回 from=变更前的状态;found=漏洞是否存在;notified=事件是否登记成功。
//
// 三条刻意的行为:
// - 状态未实际变化时不登记事件。前端抽屉重复提交同一个值、或自动化脚本
// 幂等重放,都不该产出推送噪音。
// - 漏洞不存在时返回 found=false 且不做任何写入,由调用方翻译成 404。
// - 事件登记失败不影响状态更新(见 RecordNotificationEventTx 的保存点说明),
// 所以 notified=false 时状态已经改成功了,调用方不应因此报错。
func (d *DB) SetFindingStatusWithNotify(ctx context.Context, id int64, status string) (from string, found bool, notified bool, err error) {
err = d.WithEvidenceTx(ctx, func(tx *sql.Tx) error {
var txErr error
from, found, _, notified, txErr = SetFindingStatusTx(ctx, tx, id, status)
return txErr
})
return from, found, notified, err
}
// SetFindingStatusTx 在**调用方的事务**内更新漏洞状态并登记状态变更推送事件。
//
// 抽成事务级函数是为了让所有改状态的路径共用同一套语义——此前只有
// patchFinding 走带通知的版本,而**复测结论为「已修复」时**(finding_retests
// 里那条 `UPDATE findings SET status=...`)是直接写库的,于是配了
// `on_status_change` 的渠道对这类状态流转完全收不到推送:界面上状态悄悄变了,
// 运维要到打开平台才发现。
//
// 返回 from=变更前状态、found=漏洞是否存在、changed=状态是否真的变了、
// notified=事件是否登记成功(登记失败不影响状态更新,见 RecordNotificationEventTx)。
func SetFindingStatusTx(ctx context.Context, tx *sql.Tx, id int64, status string) (from string, found bool, changed bool, notified bool, err error) {
var (
vulnclass, name, severity, summary string
taskID sql.NullInt64
assetIDs []byte
)
scanErr := tx.QueryRowContext(ctx, `SELECT vulnclass, name, severity, summary, task_id, asset_ids, status
FROM findings WHERE id=$1 FOR UPDATE`, id).
Scan(&vulnclass, &name, &severity, &summary, &taskID, &assetIDs, &from)
if scanErr == sql.ErrNoRows {
return "", false, false, false, nil
}
if scanErr != nil {
return "", false, false, false, scanErr
}
found = true
if from == status {
// 状态没有真的变化就不登记事件:重复提交同一个值、幂等重放都不该
// 产生推送噪音。
return from, true, false, false, nil
}
if _, err := tx.ExecContext(ctx, `UPDATE findings SET status=$2 WHERE id=$1`, id, status); err != nil {
return from, true, false, false, err
}
var assets []int64
_ = json.Unmarshal(assetIDs, &assets)
notified = RecordNotificationEventTx(ctx, tx, notify.EventFindingStatusChanged, id, notify.Snapshot{
Kind: notify.EventFindingStatusChanged,
FindingID: id,
TaskID: taskID.Int64,
VulnClass: vulnclass,
Name: name,
Severity: severity,
Summary: summary,
AssetIDs: assets,
FromStatus: from,
ToStatus: status,
})
return from, true, true, notified, nil
}
// NotificationStats 是通知页顶部的概览计数。
type NotificationStats struct {
Channels int `json:"channels"`
ChannelsOn int `json:"channels_on"`
Pending int `json:"pending"`
Failed int `json:"failed"`
SentToday int `json:"sent_today"`
BacklogAgeMS int64 `json:"backlog_age_ms"` // 最老的待发投递距今毫秒数
}
// NotificationStatsSnapshot 汇总通知系统的健康度。
// BacklogAgeMS 是「推送是不是卡住了」最直接的指标——比 pending 计数有用得多,
// 因为积压 3 条和积压 3 条的差别可以是从 3 秒到 3 小时。
func (d *DB) NotificationStatsSnapshot(ctx context.Context) (*NotificationStats, error) {
var s NotificationStats
if err := d.QueryRowContext(ctx, `SELECT
(SELECT count(*) FROM notification_channels),
(SELECT count(*) FROM notification_channels WHERE enabled),
(SELECT count(*) FROM notification_deliveries WHERE state IN ($1,$2)),
(SELECT count(*) FROM notification_deliveries WHERE state=$3),
(SELECT count(*) FROM notification_deliveries WHERE state=$4 AND sent_at >= date_trunc('day', now())),
COALESCE((SELECT EXTRACT(EPOCH FROM (now() - min(created_at))) * 1000 FROM notification_deliveries WHERE state=$1), 0)::bigint`,
NotifyStatePending, NotifyStateSending, NotifyStateFailed, NotifyStateSent).
Scan(&s.Channels, &s.ChannelsOn, &s.Pending, &s.Failed, &s.SentToday, &s.BacklogAgeMS); err != nil {
return nil, err
}
return &s, nil
}
+446
View File
@@ -0,0 +1,446 @@
package db
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"strings"
"time"
)
// 本文件是投递任务的领取与状态流转。
//
// 领取用「租约」而非长事务:把行置为 sending 并把 next_attempt_at 推到未来作为
// 租约到期时间,提交事务后再去做网络投递。这样投递期间不持有数据库锁——
// 网络请求可能耗时数秒(客户端超时 15 秒),占着行锁不放会拖垮同库的其它写操作。
//
// 代价是进程若在投递途中崩溃,行会停在 sending。这是**可自愈**的:租约到期后
// next_attempt_at 落入过去,下一轮领取会把同一行重新捞起来(见领取条件里的
// state IN ('pending','sending'))。重试计数在领取时就已 +1,所以崩溃不会造成
// 无限重试——MaxNotifyAttempts 次机会用完后落入 failed 等人工处理。
// MaxNotifyAttempts 是一条投递的最大尝试次数(含首次)。
// 定义在这里而非投递引擎里:它是状态机自身的策略,引擎只是执行者。
const MaxNotifyAttempts = 3
// MaxDigestBatchSize 是单个汇总批次一次最多合并多少条投递。
//
// 存在的理由是资源:一个汇总周期内如果扫出几万个漏洞(完全可能——一次全量扫描
// 就能做到),不设上界的话领取会把全部行读进内存、渲染成一条超长消息,
// 然后被渠道的长度上限截掉大半——既浪费内存,又**静默丢失**被截掉的那些漏洞。
// 设上界后,超出的部分留在库里成为下一个批次,下个周期自然发出去,不会丢。
//
// 取 500 的依据:它是渲染成消息后在企微 4096 字节上限内还"有内容可读"的量级;
// 再大也只是让截断发生在更靠后的位置而已。
const MaxDigestBatchSize = 500
// NotificationDelivery 是一条投递任务,含渲染所需的渠道配置与事件快照。
type NotificationDelivery struct {
ID int64 `json:"id"`
EventID int64 `json:"event_id"`
ChannelID int64 `json:"channel_id"`
State string `json:"state"`
Attempts int `json:"attempts"`
NextAttemptAt time.Time `json:"next_attempt_at"`
LastError string `json:"last_error"`
BatchID *int64 `json:"batch_id,omitempty"`
CreatedAt time.Time `json:"created_at"`
SentAt *time.Time `json:"sent_at,omitempty"`
Snapshot json.RawMessage `json:"snapshot,omitempty"`
// 联合加载的渲染上下文,不进 JSON(由 server 层组装 DTO)。
Channel *NotificationChannel `json:"-"`
// FindingID/EventKind 从事件带出,供历史列表直接跳转漏洞详情。
FindingID int64 `json:"finding_id,string"`
EventKind string `json:"event_kind"`
// ChannelName/ChannelKind 是列表展示用的冗余字段,省掉前端二次查询。
ChannelName string `json:"channel_name"`
ChannelKind string `json:"channel_kind"`
}
const notificationDeliveryCols = `d.id, d.event_id, d.channel_id, d.state, d.attempts, d.next_attempt_at,
d.last_error, d.batch_id, d.created_at, d.sent_at`
// joinedDeliveryQuery 是投递行的统一读取形状:投递 + 事件快照 + 渠道配置。
// 渲染一条消息三者缺一不可,分开查会写出三次往返。
const joinedDeliveryQuery = `SELECT ` + notificationDeliveryCols + `,
e.snapshot, e.kind, e.finding_id,
c.id, c.name, c.kind, c.enabled, c.config, c.mode, c.filter, c.rate_per_min
FROM notification_deliveries d
JOIN notification_events e ON e.id = d.event_id
JOIN notification_channels c ON c.id = d.channel_id`
func scanNotificationDelivery(sc interface{ Scan(...any) error }) (*NotificationDelivery, error) {
var (
dl NotificationDelivery
lastErr sql.NullString
batchID sql.NullInt64
sentAt sql.NullTime
snapshot []byte
eventKind string
channel NotificationChannel
chEnabled bool
)
if err := sc.Scan(&dl.ID, &dl.EventID, &dl.ChannelID, &dl.State, &dl.Attempts, &dl.NextAttemptAt,
&lastErr, &batchID, &dl.CreatedAt, &sentAt,
&snapshot, &eventKind, &dl.FindingID,
&channel.ID, &channel.Name, &channel.Kind, &chEnabled, &channel.Config, &channel.Mode, &channel.Filter, &channel.RatePerMin); err != nil {
return nil, err
}
dl.LastError = lastErr.String
if batchID.Valid {
dl.BatchID = &batchID.Int64
}
if sentAt.Valid {
dl.SentAt = &sentAt.Time
}
dl.Snapshot = json.RawMessage(snapshot)
dl.EventKind = eventKind
dl.ChannelName = channel.Name
dl.ChannelKind = channel.Kind
channel.Enabled = &chEnabled
dl.Channel = &channel
return &dl, nil
}
// claimQuery 描述一次领取:先按 sel 选出候选并加锁,再把它们置为 sending 并
// 延长租约。sel 里的 lease 位置由调用方用 $n 占位并自行传参。
type claimQuery struct {
sql string
args []any
}
// ClaimRealtimeDeliveries 领取某渠道一批到期的实时投递,最多 limit 条。
//
// 刻意按**单个渠道**领取而不是「全局领一批再挑着发」:限流闸在投递引擎里按渠道
// 维护,只有先知道这个渠道这一轮还能发几条、再去领同样多的行,限流才不会消耗
// 重试次数。若反过来先领后弃,被限流挡下的行已经被计过一次 attempts,
// 3 次预算会被纯粹的等待耗光,最后落进 failed。
//
// 条件含「租约已过期的 sending」——那是崩溃自愈的落点。lease 必须显著大于单次
// 投递的最坏耗时(渠道 HTTP 客户端超时 15 秒),否则同一行会被两个 dispatcher
// 同时投递。同时挡掉已停用渠道:停用操作已把存量投递标记为 skipped,
// 这里再拦一道,避免停用与领取并发时的漏网。
func (d *DB) ClaimRealtimeDeliveries(ctx context.Context, channelID int64, limit int, lease time.Duration) ([]*NotificationDelivery, error) {
if limit <= 0 {
return nil, nil
}
return d.claimDeliveries(ctx, lease, claimQuery{
sql: `SELECT dd.id FROM notification_deliveries dd
JOIN notification_channels c ON c.id = dd.channel_id
WHERE dd.channel_id = $1 AND dd.state IN ($2,$3) AND dd.next_attempt_at <= now()
AND c.enabled AND c.mode = $4
ORDER BY dd.next_attempt_at, dd.id
FOR UPDATE OF dd SKIP LOCKED
LIMIT $5`,
args: []any{channelID, NotifyStatePending, NotifyStateSending, NotifyModeRealtime, limit},
}, nil)
}
// DigestBatchDue 报告该渠道是否已攒够一个到期批次:存在待发投递,且**最老的那条**
// 年龄已达到汇总周期。
//
// 判定依据是最老投递的年龄而非墙上时钟:这样刚建好的渠道不会因为对齐到整点而
// 立刻吐出一条只有一条的「汇总」,积压很久的批次也不会再白等一轮。
//
// 与 ClaimDigestBatch 分开是因为语义不同:本函数只回答「该不该发」,
// 而领取要拿走该渠道**全部**待发行(包括尚未满年龄的那些)——否则一个周期
// 会被拆成多条消息,汇总就失去意义了。
func (d *DB) DigestBatchDue(ctx context.Context, channelID int64, minAge time.Duration) (bool, error) {
var due bool
err := d.QueryRowContext(ctx, `SELECT EXISTS (
SELECT 1 FROM notification_deliveries d
JOIN notification_channels c ON c.id = d.channel_id
WHERE d.channel_id = $1 AND d.state IN ($2,$3) AND c.enabled
GROUP BY d.channel_id
HAVING min(d.created_at) <= now() - make_interval(secs => $4)
)`, channelID, NotifyStatePending, NotifyStateSending, int64(minAge.Seconds())).Scan(&due)
return due, err
}
// ClaimDigestBatch 领取某渠道当前到期的待发投递,作为一个汇总批次,
// 单批最多 MaxDigestBatchSize 条。
//
// 同批次的所有投递共享 batch_id,用集合里的最小 id 作批次号(稳定、可读、
// 无需额外序列)。重试时用 COALESCE 保留原批次号,使「这批 N 条是一起发的」
// 在多次重试后依然成立。
//
// 按 id 升序取前 N 条而非随机取:最早产生的投递最先发出去,积压时不会出现
// 「新漏洞先发、老漏洞永远排在后面」的饥饿。
func (d *DB) ClaimDigestBatch(ctx context.Context, channelID int64, limit int, lease time.Duration) ([]*NotificationDelivery, error) {
if limit <= 0 {
return nil, nil
}
// limit 是**内存上界**,调用方传 MaxDigestBatchSize;这里再夹一道,
// 防止调用方传进一个更大的值。
//
// 刻意不接受「限流额度」充当批次大小:限流的单位是消息条数——一个批次只发
// 一条消息、消耗一个令牌,由 server 层的 takeTokens 扣除——与「一批装几条
// 漏洞」是两个不同的量纲。曾经为了让 rate_per_min 对 digest 生效而把每轮
// 请求预算传进来当批次大小,结果 rate=20/min 的渠道每批只装 1 条漏洞,
// digest 退化成带汇总文案的实时推送。要改限流请改 takeTokens 的 want,
// 不要动这里。
if limit > MaxDigestBatchSize {
limit = MaxDigestBatchSize
}
out, err := d.claimDeliveries(ctx, lease, claimQuery{
sql: `SELECT dd.id FROM notification_deliveries dd
JOIN notification_channels c ON c.id = dd.channel_id
WHERE dd.channel_id = $1 AND dd.state IN ($2,$3) AND dd.next_attempt_at <= now() AND c.enabled
ORDER BY dd.id
FOR UPDATE OF dd SKIP LOCKED
LIMIT $4`,
args: []any{channelID, NotifyStatePending, NotifyStateSending, limit},
}, func(tx *sql.Tx, ids []int64) error {
batchID := ids[0]
for _, id := range ids {
if id < batchID {
batchID = id
}
}
ph, idArgs := placeholders(2, ids)
_, err := tx.ExecContext(ctx, `UPDATE notification_deliveries SET batch_id = COALESCE(batch_id, $1)
WHERE id IN (`+ph+`)`, append([]any{batchID}, idArgs...)...)
return err
})
return out, err
}
// claimDeliveries 执行「选取 + 置 sending 延长租约 + 读取完整行」,全在一个事务里。
// postClaim 是可选的附加步骤(汇总批次用它写入 batch_id)。
func (d *DB) claimDeliveries(ctx context.Context, lease time.Duration, cq claimQuery, postClaim func(*sql.Tx, []int64) error) ([]*NotificationDelivery, error) {
tx, err := d.BeginTx(ctx, nil)
if err != nil {
return nil, err
}
defer tx.Rollback() //nolint:errcheck // 提交成功后是 no-op
ids, err := selectForClaim(ctx, tx, cq.sql, cq.args...)
if err != nil {
return nil, err
}
if len(ids) == 0 {
return nil, tx.Commit()
}
// 置 sending 并把 next_attempt_at 推到未来:这个未来时刻即租约到期时间,
// 「租约未到期」与「未到重试时间」因此共用同一个条件表达,不需要新增列。
ph, idArgs := placeholders(3, ids)
if _, err := tx.ExecContext(ctx, `UPDATE notification_deliveries
SET state=$1, attempts=attempts+1, next_attempt_at=now()+make_interval(secs => $2)
WHERE id IN (`+ph+`)`,
append([]any{NotifyStateSending, lease.Seconds()}, idArgs...)...); err != nil {
return nil, err
}
if postClaim != nil {
if err := postClaim(tx, ids); err != nil {
return nil, err
}
}
out, err := loadDeliveriesTx(ctx, tx, ids)
if err != nil {
return nil, err
}
return out, tx.Commit()
}
func selectForClaim(ctx context.Context, tx *sql.Tx, query string, args ...any) ([]int64, error) {
rows, err := tx.QueryContext(ctx, query, args...)
if err != nil {
return nil, err
}
defer rows.Close()
var ids []int64
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
return nil, err
}
ids = append(ids, id)
}
return ids, rows.Err()
}
func loadDeliveriesTx(ctx context.Context, tx *sql.Tx, ids []int64) ([]*NotificationDelivery, error) {
ph, args := placeholders(1, ids)
rows, err := tx.QueryContext(ctx, joinedDeliveryQuery+` WHERE d.id IN (`+ph+`) ORDER BY d.id`, args...)
if err != nil {
return nil, err
}
defer rows.Close()
out := []*NotificationDelivery{}
for rows.Next() {
dl, err := scanNotificationDelivery(rows)
if err != nil {
return nil, err
}
out = append(out, dl)
}
return out, rows.Err()
}
// MarkDeliveriesSent 把一批投递标记为已送达。
func (d *DB) MarkDeliveriesSent(ctx context.Context, ids []int64) error {
ph, args := placeholders(2, ids)
if len(args) == 0 {
return nil
}
_, err := d.ExecContext(ctx, `UPDATE notification_deliveries
SET state=$1, sent_at=now(), last_error='' WHERE id IN (`+ph+`)`, append([]any{NotifyStateSent}, args...)...)
return err
}
// RescheduleDeliveries 把一批投递退回 pending 并推后重试时间。
//
// 退回 pending 而不是引入新的中间状态,是为了让「还剩几次机会」只由一个地方
// 表达(MaxNotifyAttempts),避免状态机的分支随重试策略膨胀。
func (d *DB) RescheduleDeliveries(ctx context.Context, ids []int64, delay time.Duration, errMsg string) error {
ph, args := placeholders(4, ids)
if len(args) == 0 {
return nil
}
_, err := d.ExecContext(ctx, `UPDATE notification_deliveries
SET state=$1, next_attempt_at=now()+make_interval(secs => $2), last_error=$3
WHERE id IN (`+ph+`)`,
append([]any{NotifyStatePending, delay.Seconds(), truncateNotifyError(errMsg)}, args...)...)
return err
}
// DeferDeliveries 把一批投递退回 pending、立即可再领,并**撤销领取时计的那一次尝试**。
//
// 用途只有一个:汇总消息按渠道长度上限分段发送时,没装进本条的条目要留到下一批。
// 那不是失败,所以不该消耗重试预算——领取时 attempts 已经乐观地 +1 了,
// 这里必须减回去。否则一个 500 条的积压会按每段 20 条切成 25 段,
// 尾部条目在第 3 段就被 MaxNotifyAttempts 判成 failed,而它们从未出过任何错。
//
// GREATEST(...,0) 兜住「有人手工重发把 attempts 清零后又走到这里」的情况,
// 不让计数变成负数。
func (d *DB) DeferDeliveries(ctx context.Context, ids []int64, reason string) error {
ph, args := placeholders(3, ids)
if len(args) == 0 {
return nil
}
_, err := d.ExecContext(ctx, `UPDATE notification_deliveries
SET state=$1, attempts=GREATEST(attempts-1, 0), next_attempt_at=now(), last_error=$2
WHERE id IN (`+ph+`)`,
append([]any{NotifyStatePending, truncateNotifyError(reason)}, args...)...)
return err
}
// FailDeliveries 把一批投递标记为最终失败,等待人工在投递历史里重发。
func (d *DB) FailDeliveries(ctx context.Context, ids []int64, errMsg string) error {
// 占位符从 $3 开始:$1 是 state、$2 是 last_error。
ph, args := placeholders(3, ids)
if len(args) == 0 {
return nil
}
_, err := d.ExecContext(ctx, `UPDATE notification_deliveries SET state=$1, last_error=$2 WHERE id IN (`+ph+`)`,
append([]any{NotifyStateFailed, truncateNotifyError(errMsg)}, args...)...)
return err
}
// RetryNotificationDelivery 手动重发一条投递:重置为 pending、清零重试计数、
// 立即到期。清计数是刻意的——人工点「重发」意味着前几次失败的原因已被处理,
// 再拿旧计数限制它没有道理。
func (d *DB) RetryNotificationDelivery(ctx context.Context, id int64) error {
res, err := d.ExecContext(ctx, `UPDATE notification_deliveries
SET state=$2, attempts=0, next_attempt_at=now(), last_error=''
WHERE id=$1 AND state IN ($3,$4)`, id, NotifyStatePending, NotifyStateFailed, NotifyStateSkipped)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return fmt.Errorf("投递 %d 不存在或当前状态不允许重发", id)
}
return nil
}
// NotificationDeliveryFilter 是投递历史的查询条件。
type NotificationDeliveryFilter struct {
ChannelID int64
State string
EventKind string
}
func (f NotificationDeliveryFilter) where() (string, []any) {
var conds []string
var args []any
if f.ChannelID > 0 {
args = append(args, f.ChannelID)
conds = append(conds, fmt.Sprintf("d.channel_id=$%d", len(args)))
}
if f.State != "" {
args = append(args, f.State)
conds = append(conds, fmt.Sprintf("d.state=$%d", len(args)))
}
if f.EventKind != "" {
args = append(args, f.EventKind)
conds = append(conds, fmt.Sprintf("e.kind=$%d", len(args)))
}
if len(conds) == 0 {
return "", nil
}
return " WHERE " + strings.Join(conds, " AND "), args
}
// ListNotificationDeliveries 分页返回投递历史,新的在前。
func (d *DB) ListNotificationDeliveries(ctx context.Context, f NotificationDeliveryFilter, page, pageSize int) ([]*NotificationDelivery, int, error) {
if page < 1 {
page = 1
}
if pageSize <= 0 || pageSize > 200 {
pageSize = 50
}
where, args := f.where()
var total int
if err := d.QueryRowContext(ctx, `SELECT count(*) FROM notification_deliveries d
JOIN notification_events e ON e.id = d.event_id`+where, args...).Scan(&total); err != nil {
return nil, 0, err
}
q := fmt.Sprintf("%s%s ORDER BY d.id DESC LIMIT $%d OFFSET $%d",
joinedDeliveryQuery, where, len(args)+1, len(args)+2)
rows, err := d.QueryContext(ctx, q, append(args, pageSize, (page-1)*pageSize)...)
if err != nil {
return nil, 0, err
}
defer rows.Close()
out := []*NotificationDelivery{}
for rows.Next() {
dl, err := scanNotificationDelivery(rows)
if err != nil {
return nil, 0, err
}
out = append(out, dl)
}
return out, total, rows.Err()
}
// truncateNotifyError 把错误信息截到列可接受的长度。渠道返回的响应体可能很长
// (通用 Webhook 打到自建服务时尤甚),不截断会让历史列表的载荷膨胀。
func truncateNotifyError(msg string) string {
const max = 500
if len(msg) <= max {
return msg
}
// 按字符边界回退,避免留下半个 UTF-8 字符让前端显示成乱码。
cut := max
for cut > 0 && !isUTF8Start(msg[cut]) {
cut--
}
return msg[:cut] + "…"
}
func isUTF8Start(b byte) bool { return b&0xC0 != 0x80 }
// placeholders 生成从 start 开始的 $n 占位串及对应参数,供 IN (...) 使用。
// 例如 start=3, ids=[7,8] → "$3,$4", [7,8]。
func placeholders(start int, ids []int64) (string, []any) {
ph := make([]string, 0, len(ids))
args := make([]any, 0, len(ids))
for i, id := range ids {
ph = append(ph, fmt.Sprintf("$%d", start+i))
args = append(args, id)
}
return strings.Join(ph, ","), args
}
+874
View File
@@ -0,0 +1,874 @@
package db
import (
"context"
"encoding/json"
"testing"
"time"
"github.com/Autumn-27/artex/notify"
)
// 本文件的用例都会真连 PostgreSQL(无库时跳过)。这些 SQL 用到了
// FOR UPDATE SKIP LOCKED、make_interval、JSONB、多行 IN(...) 占位符拼接,
// 都是「编译通过但可能运行时报错」的写法,必须实跑才算验证过。
func notifyTestDB(t *testing.T) *DB {
t.Helper()
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
t.Cleanup(func() { d.Close() })
return d
}
// newTestChannel 建一个渠道,测试结束自动删除。
func newTestChannel(t *testing.T, d *DB, kind, mode string, filter string) *NotificationChannel {
t.Helper()
if filter == "" {
filter = `{}`
}
ch := &NotificationChannel{
Name: "测试渠道-" + t.Name(),
Kind: kind,
Mode: mode,
Config: json.RawMessage(`{"webhook":"https://example.com/hook"}`),
Filter: json.RawMessage(filter),
RatePerMin: 100,
}
id, err := d.SaveNotificationChannel(context.Background(), ch)
if err != nil {
t.Fatalf("建渠道失败: %v", err)
}
t.Cleanup(func() { d.Exec(`DELETE FROM notification_channels WHERE id=$1`, id) })
ch.ID = id
return ch
}
// addTestEvent 直接写一条事件(不经 finding),用于测试分派与投递。
func addTestEvent(t *testing.T, d *DB, kind string, findingID int64, snap notify.Snapshot) int64 {
t.Helper()
snap.Kind = kind
snap.FindingID = findingID
id, err := d.AddNotificationEvent(context.Background(), kind, findingID, snap)
if err != nil {
t.Fatalf("写事件失败: %v", err)
}
t.Cleanup(func() { d.Exec(`DELETE FROM notification_events WHERE id=$1`, id) })
return id
}
func TestNotificationAssetNamesResolvesAndPreservesOrder(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
// 三类资产各有各的展示口径:域名、IP、URL。
insertAsset := func(query, value string) int64 {
t.Helper()
var id int64
if err := d.QueryRow(query, value).Scan(&id); err != nil {
t.Fatal(err)
}
return id
}
domID := insertAsset(`INSERT INTO assets(type, domain) VALUES('subdomain',$1) RETURNING id`, "a.example.com")
ipID := insertAsset(`INSERT INTO assets(type, ip) VALUES('ip',$1) RETURNING id`, "10.1.2.3")
svcID := insertAsset(`INSERT INTO assets(type, url) VALUES('service',$1) RETURNING id`, "https://a.example.com/admin")
t.Cleanup(func() {
d.Exec(`DELETE FROM assets WHERE id IN ($1,$2,$3)`, domID, ipID, svcID)
})
// 传入顺序刻意乱序,且含一个不存在的 id。
got, err := d.NotificationAssetNames(ctx, []int64{svcID, 999999999, domID, ipID, svcID})
if err != nil {
t.Fatalf("解析资产名失败: %v", err)
}
want := []string{"https://a.example.com/admin", "a.example.com", "10.1.2.3"}
if len(got) != len(want) {
t.Fatalf("资产名数量不符,期望 %v 得到 %v", want, got)
}
for i := range want {
if got[i] != want[i] {
t.Fatalf("顺序/取值不符,期望 %v 得到 %v", want, got)
}
}
}
// TestRecordNotificationEventTxUnwindsOnFailure 是保存点机制的核心用例:
// 在事务里先让 notification_events 的写入必然失败(临时加一个恒 false 的约束),
// 断言 ① 该函数报 false ② 事务没有进入 aborted 状态,后续语句仍能执行。
//
// 没有保存点的话,PostgreSQL 会让整个事务作废,后续任何语句都以
// "current transaction is aborted" 失败——那正是「一个通知表的问题导致
// 漏洞存不进库」的故障路径。
//
// 这里刻意用 **ROLLBACK 收尾而不是 COMMIT**:ALTER TABLE 在 PG 里是事务性的,
// 一旦提交,那个临时约束就会永久留在 schema 里,把后续所有用例一起打挂。
// 回滚能自动撤销 DDL,无需手工清理。断言只需要「事务还活着」,
// 不需要真的提交。
func TestRecordNotificationEventTxUnwindsOnFailure(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
// 防御性清理:若历史运行留下过这个约束,先摘掉。
if _, err := d.Exec(`ALTER TABLE notification_events DROP CONSTRAINT IF EXISTS notify_test_never`); err != nil {
t.Fatal(err)
}
tx, err := d.BeginTx(ctx, nil)
if err != nil {
t.Fatal(err)
}
defer tx.Rollback() //nolint:errcheck // 撤销临时约束,见函数注释
// NOT VALID:只约束此后写入的行,不去校验库里已有的历史事件
// (否则存量行违规会导致约束加不上)。
if _, err := tx.ExecContext(ctx, `ALTER TABLE notification_events ADD CONSTRAINT notify_test_never CHECK (false) NOT VALID`); err != nil {
t.Fatalf("加临时约束失败: %v", err)
}
if RecordNotificationEventTx(ctx, tx, notify.EventFindingCreated, 1, notify.Snapshot{Severity: "high"}) {
t.Fatal("在必然失败的约束下仍报告写入成功")
}
// 关键断言:事务还能用。
var one int
if err := tx.QueryRowContext(ctx, `SELECT 1`).Scan(&one); err != nil {
t.Fatalf("事务已被污染(保存点未生效): %v", err)
}
if err := tx.Rollback(); err != nil {
t.Fatalf("回滚失败: %v", err)
}
// 确认 DDL 已随回滚撤销,不给后续用例留雷。
var exists bool
if err := d.QueryRow(`SELECT EXISTS(SELECT 1 FROM pg_constraint WHERE conname='notify_test_never')`).Scan(&exists); err != nil {
t.Fatal(err)
}
if exists {
t.Fatal("临时约束未被回滚撤销,会污染后续用例")
}
}
func TestFanOutRoutesEventsByFilter(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
all := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
onlyCritical := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{"min_severity":"critical"}`)
sqlOnly := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{"vulnclass_include":["SQL"]}`)
highSQL := addTestEvent(t, d, notify.EventFindingCreated, 1001, notify.Snapshot{Severity: "high", VulnClass: "SQL注入"})
lowXSS := addTestEvent(t, d, notify.EventFindingCreated, 1002, notify.Snapshot{Severity: "low", VulnClass: "XSS"})
criticalXSS := addTestEvent(t, d, notify.EventFindingCreated, 1003, notify.Snapshot{Severity: "critical", VulnClass: "XSS"})
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
t.Fatalf("分派失败: %v", err)
}
cases := []struct {
name string
eventID int64
channel int64
want bool
}{
{"全收渠道收到 high", highSQL, all.ID, true},
{"全收渠道收到 low", lowXSS, all.ID, true},
{"仅严重渠道跳过 high", highSQL, onlyCritical.ID, false},
{"仅严重渠道收到 critical", criticalXSS, onlyCritical.ID, true},
{"仅SQL渠道收到 SQL", highSQL, sqlOnly.ID, true},
{"仅SQL渠道跳过 XSS", lowXSS, sqlOnly.ID, false},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
var exists bool
if err := d.QueryRowContext(ctx, `SELECT EXISTS(SELECT 1 FROM notification_deliveries WHERE event_id=$1 AND channel_id=$2)`,
tc.eventID, tc.channel).Scan(&exists); err != nil {
t.Fatal(err)
}
if exists != tc.want {
t.Fatalf("投递是否存在: 期望 %v 得到 %v", tc.want, exists)
}
})
}
// 再分派一次不应产生重复投递(fanned_out 幂等)。
events, deliveries, err := d.FanOutPendingEvents(ctx, 100)
if err != nil {
t.Fatal(err)
}
if events != 0 || deliveries != 0 {
t.Fatalf("已分派的事件不应被再次处理,得到 events=%d deliveries=%d", events, deliveries)
}
}
// TestFanOutMarksEventsWithNoMatchingChannel 覆盖「事件没命中任何渠道」的情况。
// 这类事件必须照样被标记为已分派,否则它会永远留在待分派集合里、每个 tick 重扫。
func TestFanOutMarksEventsWithNoMatchingChannel(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
pick := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{"vulnclass_include":["绝不匹配的类型"]}`)
_ = pick
ev := addTestEvent(t, d, notify.EventFindingCreated, 2001, notify.Snapshot{Severity: "high", VulnClass: "XSS"})
_, deliveries, err := d.FanOutPendingEvents(ctx, 100)
if err != nil {
t.Fatal(err)
}
if deliveries != 0 {
t.Fatalf("不该产生投递,得到 %d", deliveries)
}
var fanned bool
if err := d.QueryRowContext(ctx, `SELECT fanned_out FROM notification_events WHERE id=$1`, ev).Scan(&fanned); err != nil {
t.Fatal(err)
}
if !fanned {
t.Fatal("未命中渠道的事件也必须标记为已分派,否则会被无限重扫")
}
}
func TestClaimRealtimeDeliveriesHonorsLeaseAndMode(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
realtime := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
digest := newTestChannel(t, d, notify.KindDingTalk, NotifyModeDigest, `{}`)
addTestEvent(t, d, notify.EventFindingCreated, 3001, notify.Snapshot{Severity: "high", VulnClass: "XSS"})
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
t.Fatal(err)
}
// 实时领取只应拿到 realtime 渠道的那条,不该动 digest 渠道的。
got, err := d.ClaimRealtimeDeliveries(ctx, realtime.ID, 10, time.Minute)
if err != nil {
t.Fatalf("领取失败: %v", err)
}
if len(got) != 1 {
t.Fatalf("应领到 1 条,得到 %d", len(got))
}
if got[0].State != NotifyStateSending || got[0].Attempts != 1 {
t.Fatalf("领取后应为 sending 且 attempts=1,得到 state=%s attempts=%d", got[0].State, got[0].Attempts)
}
// 关联加载的渲染上下文必须齐全(渠道配置 + 事件快照 + finding id)。
if got[0].Channel == nil || len(got[0].Channel.Config) == 0 {
t.Fatal("领取结果缺少渠道配置,渲染会失败")
}
if got[0].FindingID != 3001 {
t.Fatalf("finding id 未从事件带出,得到 %d", got[0].FindingID)
}
// 租约未到期,第二次领取应为空——这是「同一行不会被两个 dispatcher 同时投递」
// 的保证。
again, err := d.ClaimRealtimeDeliveries(ctx, realtime.ID, 10, time.Minute)
if err != nil {
t.Fatal(err)
}
if len(again) != 0 {
t.Fatalf("租约期内不应重复领取,得到 %d 条", len(again))
}
// digest 渠道的投递不应被实时领取碰到。
left, err := d.ClaimRealtimeDeliveries(ctx, digest.ID, 10, time.Minute)
if err != nil {
t.Fatal(err)
}
if len(left) != 0 {
t.Fatalf("实时领取不应拿到 digest 渠道的投递,得到 %d 条", len(left))
}
}
// TestClaimExpiredLeaseRecovers 覆盖崩溃自愈:进程在投递途中挂掉会留下 sending
// 行,租约到期后必须能被重新领起来,否则这条投递永远卡住。
func TestClaimExpiredLeaseRecovers(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
addTestEvent(t, d, notify.EventFindingCreated, 4001, notify.Snapshot{Severity: "high"})
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
t.Fatal(err)
}
first, err := d.ClaimRealtimeDeliveries(ctx, ch.ID, 10, time.Minute)
if err != nil || len(first) != 1 {
t.Fatalf("首次领取失败: %v (%d 条)", err, len(first))
}
// 把租约手动推到过去,模拟「租约已过期」。
if _, err := d.Exec(`UPDATE notification_deliveries SET next_attempt_at = now() - interval '1 minute' WHERE id=$1`, first[0].ID); err != nil {
t.Fatal(err)
}
second, err := d.ClaimRealtimeDeliveries(ctx, ch.ID, 10, time.Minute)
if err != nil {
t.Fatal(err)
}
if len(second) != 1 {
t.Fatalf("租约过期的 sending 行应可被重新领取,得到 %d 条", len(second))
}
if second[0].Attempts != 2 {
t.Fatalf("重新领取应累加尝试次数,得到 %d", second[0].Attempts)
}
}
func TestClaimSkipsDisabledChannel(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
addTestEvent(t, d, notify.EventFindingCreated, 5001, notify.Snapshot{Severity: "high"})
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
t.Fatal(err)
}
// 停用会把存量待发投递一起标记为 skipped。
if err := d.SetNotificationChannelEnabled(ctx, ch.ID, false); err != nil {
t.Fatal(err)
}
var state string
if err := d.QueryRow(`SELECT state FROM notification_deliveries WHERE channel_id=$1`, ch.ID).Scan(&state); err != nil {
t.Fatal(err)
}
if state != NotifyStateSkipped {
t.Fatalf("停用渠道的存量待发投递应被标记为 skipped,得到 %s", state)
}
got, err := d.ClaimRealtimeDeliveries(ctx, ch.ID, 10, time.Minute)
if err != nil {
t.Fatal(err)
}
if len(got) != 0 {
t.Fatalf("已停用渠道不应能被领取,得到 %d 条", len(got))
}
}
func TestDigestBatchDueAndStableBatchID(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeDigest, `{}`)
for i := 0; i < 3; i++ {
addTestEvent(t, d, notify.EventFindingCreated, int64(6000+i), notify.Snapshot{Severity: "high"})
}
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
t.Fatal(err)
}
// 批次刚建、年龄为 0,30 分钟的周期下不该到期。
due, err := d.DigestBatchDue(ctx, ch.ID, 30*time.Minute)
if err != nil {
t.Fatalf("判断批次到期失败: %v", err)
}
if due {
t.Fatal("刚建立的批次不应立即到期")
}
// 把三条投递的创建时间一起推老,模拟一个攒够周期的批次。
if _, err := d.Exec(`UPDATE notification_deliveries SET created_at = now() - interval '40 minutes' WHERE channel_id=$1`, ch.ID); err != nil {
t.Fatal(err)
}
due, err = d.DigestBatchDue(ctx, ch.ID, 30*time.Minute)
if err != nil {
t.Fatal(err)
}
if !due {
t.Fatal("超过周期的批次应判定为到期")
}
batch, err := d.ClaimDigestBatch(ctx, ch.ID, MaxDigestBatchSize, time.Minute)
if err != nil {
t.Fatalf("领取汇总批次失败: %v", err)
}
if len(batch) != 3 {
t.Fatalf("汇总应一次领走全部 3 条,得到 %d 条", len(batch))
}
if batch[0].BatchID == nil {
t.Fatal("汇总批次必须写 batch_id,否则历史里看不出这几条是一起发的")
}
firstBatchID := *batch[0].BatchID
for _, dl := range batch {
if dl.BatchID == nil || *dl.BatchID != firstBatchID {
t.Fatalf("同一批次应共享 batch_id,得到 %v vs %d", dl.BatchID, firstBatchID)
}
}
// 让这批**整体**失败重排后再领,batch_id 必须保持原值(COALESCE 的作用):
// 否则一次重试就把「这批是一起发的」这个事实抹掉了。
//
// 必须整批重排而不是只重排一条——投递引擎发汇总消息时就是这样处理的
// (一条消息代表整批,成败与共)。只重排一条的话,其余仍在租约期内,
// 重领自然只拿到那一条。
allIDs := make([]int64, 0, len(batch))
for _, dl := range batch {
allIDs = append(allIDs, dl.ID)
}
if err := d.RescheduleDeliveries(ctx, allIDs, time.Second, "模拟失败"); err != nil {
t.Fatal(err)
}
// 把租约推到过去,模拟退避时间已到。
if _, err := d.Exec(`UPDATE notification_deliveries SET next_attempt_at = now() - interval '1 minute' WHERE channel_id=$1`, ch.ID); err != nil {
t.Fatal(err)
}
reclaimed, err := d.ClaimDigestBatch(ctx, ch.ID, MaxDigestBatchSize, time.Minute)
if err != nil {
t.Fatal(err)
}
if len(reclaimed) != 3 {
t.Fatalf("重领应拿到全部 3 条,得到 %d", len(reclaimed))
}
if reclaimed[0].BatchID == nil || *reclaimed[0].BatchID != firstBatchID {
t.Fatalf("重试后 batch_id 应保持原值 %d,得到 %v", firstBatchID, reclaimed[0].BatchID)
}
}
func TestDeliveryStateTransitions(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
addTestEvent(t, d, notify.EventFindingCreated, 7001, notify.Snapshot{Severity: "high"})
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
t.Fatal(err)
}
got, err := d.ClaimRealtimeDeliveries(ctx, ch.ID, 10, time.Minute)
if err != nil || len(got) != 1 {
t.Fatalf("领取失败: %v (%d)", err, len(got))
}
id := got[0].ID
if err := d.RescheduleDeliveries(ctx, []int64{id}, time.Second, "网络抖动"); err != nil {
t.Fatal(err)
}
var state, lastErr string
if err := d.QueryRow(`SELECT state, last_error FROM notification_deliveries WHERE id=$1`, id).Scan(&state, &lastErr); err != nil {
t.Fatal(err)
}
if state != NotifyStatePending || lastErr != "网络抖动" {
t.Fatalf("重排后应为 pending 并记录原因,得到 state=%s err=%q", state, lastErr)
}
if err := d.FailDeliveries(ctx, []int64{id}, "重试耗尽"); err != nil {
t.Fatal(err)
}
if err := d.QueryRow(`SELECT state FROM notification_deliveries WHERE id=$1`, id).Scan(&state); err != nil {
t.Fatal(err)
}
if state != NotifyStateFailed {
t.Fatalf("应为 failed,得到 %s", state)
}
// 手动重发要清零重试计数并立即到期,否则会继承旧的失败预算。
if err := d.RetryNotificationDelivery(ctx, id); err != nil {
t.Fatalf("重发失败: %v", err)
}
var attempts int
var next time.Time
if err := d.QueryRow(`SELECT state, attempts, next_attempt_at FROM notification_deliveries WHERE id=$1`, id).Scan(&state, &attempts, &next); err != nil {
t.Fatal(err)
}
if state != NotifyStatePending || attempts != 0 {
t.Fatalf("重发后应为 pending 且 attempts=0,得到 state=%s attempts=%d", state, attempts)
}
if next.After(time.Now().Add(time.Second)) {
t.Fatal("重发应立即可领")
}
// 已送达的投递不应能被重发。
if err := d.MarkDeliveriesSent(ctx, []int64{id}); err != nil {
t.Fatal(err)
}
if err := d.RetryNotificationDelivery(ctx, id); err == nil {
t.Fatal("已送达的投递不该允许重发")
}
}
func TestListNotificationDeliveriesPagingAndFilter(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
for i := 0; i < 5; i++ {
addTestEvent(t, d, notify.EventFindingCreated, int64(8000+i), notify.Snapshot{Severity: "high", Name: "分页测试"})
}
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
t.Fatal(err)
}
if _, err := d.ClaimRealtimeDeliveries(ctx, ch.ID, 10, time.Minute); err != nil {
t.Fatal(err)
}
page1, total, err := d.ListNotificationDeliveries(ctx, NotificationDeliveryFilter{ChannelID: ch.ID, State: NotifyStateSending}, 1, 2)
if err != nil {
t.Fatalf("查询失败: %v", err)
}
if total != 5 {
t.Fatalf("总数应为 5,得到 %d", total)
}
if len(page1) != 2 {
t.Fatalf("每页 2 条,得到 %d", len(page1))
}
// 新的在前:第一页首条 id 应大于第二页首条。
page2, _, err := d.ListNotificationDeliveries(ctx, NotificationDeliveryFilter{ChannelID: ch.ID, State: NotifyStateSending}, 2, 2)
if err != nil {
t.Fatal(err)
}
if len(page2) != 2 || page2[0].ID >= page1[0].ID {
t.Fatalf("分页顺序应为新的在前,得到 page1[0]=%d page2[0]=%d", page1[0].ID, page2[0].ID)
}
// 渲染上下文必须随历史一起返回,否则列表无法展示「推的是什么」。
if page1[0].ChannelName == "" || page1[0].FindingID == 0 {
t.Fatalf("历史项缺少展示字段: %+v", page1[0])
}
// 按状态过滤:没有 pending 的。
pending, totalPending, err := d.ListNotificationDeliveries(ctx, NotificationDeliveryFilter{ChannelID: ch.ID, State: NotifyStatePending}, 1, 50)
if err != nil {
t.Fatal(err)
}
if totalPending != 0 || len(pending) != 0 {
t.Fatalf("不该有 pending 投递,得到 %d 条 (total=%d)", len(pending), totalPending)
}
}
func TestSetFindingStatusWithNotifyOnlyEmitsOnRealChange(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
tk, err := d.CreateTask("通知状态变更测试", "目标", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
defer d.DeleteTask(tk.ID)
es := d.Exploration(tk.ExplorationID)
f, err := es.RecordFinding(ctx, RecordFindingInput{
TaskID: tk.ID, Worker: "test", VulnClass: "SQL注入", Name: "状态变更用例",
Severity: "high", Summary: "摘要",
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { d.Exec(`DELETE FROM notification_events WHERE finding_id=$1`, f.FindingID) })
// 落库时已登记一条 finding_created 事件,先把它数出来作为基线。
var base int
if err := d.QueryRow(`SELECT count(*) FROM notification_events WHERE finding_id=$1`, f.FindingID).Scan(&base); err != nil {
t.Fatal(err)
}
if base < 1 {
t.Fatal("漏洞落库应在同一事务里登记一条推送事件")
}
// 改成同一个状态:不该产生事件(避免重复提交刷出推送噪音)。
from, found, notified, err := d.SetFindingStatusWithNotify(ctx, f.FindingID, "pending")
if err != nil || !found {
t.Fatalf("状态设置失败: found=%v err=%v", found, err)
}
if notified {
t.Fatal("状态未变化时不应登记推送事件")
}
if from != "pending" {
t.Fatalf("应返回变更前状态 pending,得到 %q", from)
}
// 真正变更:应登记事件并记录 from/to。
from, found, notified, err = d.SetFindingStatusWithNotify(ctx, f.FindingID, "fixed")
if err != nil || !found {
t.Fatalf("状态设置失败: found=%v err=%v", found, err)
}
if !notified {
t.Fatal("状态实际变更时应登记推送事件")
}
if from != "pending" {
t.Fatalf("from 应为 pending,得到 %q", from)
}
var snapshot []byte
if err := d.QueryRow(`SELECT snapshot FROM notification_events WHERE finding_id=$1 AND kind=$2`,
f.FindingID, notify.EventFindingStatusChanged).Scan(&snapshot); err != nil {
t.Fatalf("未找到状态变更事件: %v", err)
}
var snap notify.Snapshot
if err := json.Unmarshal(snapshot, &snap); err != nil {
t.Fatal(err)
}
if snap.FromStatus != "pending" || snap.ToStatus != "fixed" {
t.Fatalf("快照里的状态流转不对: %s → %s", snap.FromStatus, snap.ToStatus)
}
// 快照要带上渲染所需字段,否则状态变更消息会是空壳。
if snap.VulnClass != "SQL注入" || snap.Severity != "high" || snap.Name != "状态变更用例" {
t.Fatalf("快照缺少渲染字段: %+v", snap)
}
var status string
if err := d.QueryRow(`SELECT status FROM findings WHERE id=$1`, f.FindingID).Scan(&status); err != nil {
t.Fatal(err)
}
if status != "fixed" {
t.Fatalf("状态应已更新为 fixed,得到 %s", status)
}
// 不存在的漏洞:found=false,不报错。
if _, found, _, err := d.SetFindingStatusWithNotify(ctx, 999999999, "fixed"); err != nil || found {
t.Fatalf("不存在的漏洞应返回 found=false 且无错,得到 found=%v err=%v", found, err)
}
}
func TestNotificationStatsSnapshot(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
addTestEvent(t, d, notify.EventFindingCreated, 9001, notify.Snapshot{Severity: "high"})
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
t.Fatal(err)
}
stats, err := d.NotificationStatsSnapshot(ctx)
if err != nil {
t.Fatalf("统计失败: %v", err)
}
if stats.Channels < 1 || stats.ChannelsOn < 1 {
t.Fatalf("渠道计数不对: %+v", stats)
}
if stats.Pending < 1 {
t.Fatalf("应统计到待发投递: %+v", stats)
}
// 刚建的投递积压年龄应接近 0,而不是负数或巨大值。
if stats.BacklogAgeMS < 0 || stats.BacklogAgeMS > int64(time.Hour/time.Millisecond) {
t.Fatalf("积压年龄不合法: %d ms", stats.BacklogAgeMS)
}
_ = ch
}
func TestNotificationChannelCRUDRoundTrip(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
ch := &NotificationChannel{
Name: "CRUD 往返",
Kind: notify.KindEmail,
Mode: NotifyModeDigest,
Config: json.RawMessage(`{"host":"smtp.example.com","port":587,"from":"a@b.c","to":["x@y.z"]}`),
Filter: json.RawMessage(`{"min_severity":"medium","on_status_change":true}`),
RatePerMin: 42,
}
id, err := d.SaveNotificationChannel(ctx, ch)
if err != nil {
t.Fatalf("新建失败: %v", err)
}
t.Cleanup(func() { d.Exec(`DELETE FROM notification_channels WHERE id=$1`, id) })
got, err := d.NotificationChannelByID(ctx, id)
if err != nil {
t.Fatalf("读取失败: %v", err)
}
if got.Mode != NotifyModeDigest || got.RatePerMin != 42 || got.Name != "CRUD 往返" {
t.Fatalf("往返字段不一致: %+v", got)
}
if !got.IsEnabled() {
t.Fatal("默认应为启用")
}
var cfg map[string]any
if err := json.Unmarshal(got.Config, &cfg); err != nil {
t.Fatal(err)
}
if cfg["host"] != "smtp.example.com" {
t.Fatalf("配置未正确落库: %v", cfg)
}
var filter notify.Filter
if err := json.Unmarshal(got.Filter, &filter); err != nil {
t.Fatal(err)
}
if filter.MinSeverity != "medium" || !filter.OnStatusChange {
t.Fatalf("过滤条件未正确落库: %+v", filter)
}
// 更新后再读。
got.Name = "改名了"
off := false
got.Enabled = &off
if _, err := d.SaveNotificationChannel(ctx, got); err != nil {
t.Fatal(err)
}
after, err := d.NotificationChannelByID(ctx, id)
if err != nil {
t.Fatal(err)
}
if after.Name != "改名了" || after.IsEnabled() {
t.Fatalf("更新未生效: %+v", after)
}
// 删除后应报「不存在」而不是静默成功。
if err := d.DeleteNotificationChannel(ctx, id); err != nil {
t.Fatal(err)
}
if _, err := d.NotificationChannelByID(ctx, id); err != ErrNotificationChannelNotFound {
t.Fatalf("期望 ErrNotificationChannelNotFound,得到 %v", err)
}
if err := d.DeleteNotificationChannel(ctx, id); err != ErrNotificationChannelNotFound {
t.Fatalf("重复删除应报不存在,得到 %v", err)
}
}
// TestSaveNotificationChannelKeepsExplicitZeroRate 锁住一个曾经写错的地方:
// **0 是合法配置,含义是「不限流」,不能被 db 层当成「未指定」覆盖成默认值**。
//
// 历史 bug:SaveNotificationChannel 里写了 `if RatePerMin <= 0 { 取默认值 }`,
// 于是文档、UI 提示、takeTokens 都按「0=不限流」解释,唯独写库这一层悄悄改成
// 20(钉钉/企微/Telegram)或 100(飞书)——操作者以为放开了限流、实际被卡着,
// 而且没有任何提示。「未指定」与「显式 0」的区别只有请求体能表达,
// 所以默认值在 server 层填(见 notifyCreateChannel),db 层只管存。
func TestSaveNotificationChannelKeepsExplicitZeroRate(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
// 显式 0(不限流):必须原样存下来。
unlimited := &NotificationChannel{
Name: "不限流", Kind: notify.KindDingTalk, RatePerMin: 0,
Config: json.RawMessage(`{"webhook":"https://example.com/h"}`),
}
id, err := d.SaveNotificationChannel(ctx, unlimited)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { d.Exec(`DELETE FROM notification_channels WHERE id=$1`, id) })
got, err := d.NotificationChannelByID(ctx, id)
if err != nil {
t.Fatal(err)
}
if got.RatePerMin != 0 {
t.Fatalf("显式 0 表示不限流,必须原样保存,得到 %d", got.RatePerMin)
}
if got.Mode != NotifyModeRealtime {
t.Fatalf("默认模式应为 realtime,得到 %s", got.Mode)
}
// 负值是非法输入,应被拒绝而不是悄悄改成别的值。
bad := &NotificationChannel{
Name: "负限流", Kind: notify.KindDingTalk, RatePerMin: -1,
Config: json.RawMessage(`{"webhook":"https://example.com/h"}`),
}
if _, err := d.SaveNotificationChannel(ctx, bad); err == nil {
t.Fatal("负限流应被拒绝")
}
}
// TestDeleteChannelCascadesDeliveries 锁住外键行为:渠道删除后其投递历史一并消失
// (配置都没了,历史无从解读),但事件本身要留下——它可能还被别的渠道引用。
func TestDeleteChannelCascadesDeliveries(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
ev := addTestEvent(t, d, notify.EventFindingCreated, 9101, notify.Snapshot{Severity: "high"})
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
t.Fatal(err)
}
var before int
if err := d.QueryRow(`SELECT count(*) FROM notification_deliveries WHERE channel_id=$1`, ch.ID).Scan(&before); err != nil {
t.Fatal(err)
}
if before == 0 {
t.Fatal("前置条件不成立:未产生投递")
}
if err := d.DeleteNotificationChannel(ctx, ch.ID); err != nil {
t.Fatal(err)
}
var after int
if err := d.QueryRow(`SELECT count(*) FROM notification_deliveries WHERE channel_id=$1`, ch.ID).Scan(&after); err != nil {
t.Fatal(err)
}
if after != 0 {
t.Fatalf("渠道删除后其投递应级联删除,仍有 %d 条", after)
}
var evExists bool
if err := d.QueryRow(`SELECT EXISTS(SELECT 1 FROM notification_events WHERE id=$1)`, ev).Scan(&evExists); err != nil {
t.Fatal(err)
}
if !evExists {
t.Fatal("删渠道不应连带删除事件本身")
}
}
// TestClaimDigestBatchHonorsCallerLimit 覆盖审计指出的一处口子:
// 汇总渠道此前完全绕过令牌桶——allow 被 takeTokens 扣掉却没人用,
// rate_per_min 对 digest 模式毫无作用。现在 limit 也参与约束。
func TestClaimDigestBatchHonorsCallerLimit(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeDigest, `{}`)
for i := 0; i < 10; i++ {
addTestEvent(t, d, notify.EventFindingCreated, int64(7000+i), notify.Snapshot{Severity: "high"})
}
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
t.Fatal(err)
}
// 取 limit=3:只能领到 3 条,其余留在库里。
got, err := d.ClaimDigestBatch(ctx, ch.ID, 3, time.Minute)
if err != nil {
t.Fatalf("领取失败: %v", err)
}
if len(got) != 3 {
t.Fatalf("应按调用方限流额度只领 3 条,得到 %d", len(got))
}
// limit=0 表示本轮额度用尽:一条都不该领,也不该报错。
if got, err := d.ClaimDigestBatch(ctx, ch.ID, 0, time.Minute); err != nil || len(got) != 0 {
t.Fatalf("额度为 0 时应领 0 条且不报错,得到 %d 条 err=%v", len(got), err)
}
}
// TestFinishFindingRetestEmitsStatusChange 覆盖审计指出的一处完整性缺口:
// 复测结论为「已修复」时,状态确实变了,但那条 UPDATE 是直接写库的、
// 绕过了带通知的版本——于是配了 on_status_change 的渠道对这种状态流转
// 完全收不到推送,界面上状态悄悄变了,运维要打开平台才知道。
//
// 这条用例锁住「所有改状态的路径都要登记状态变更事件」。
func TestFinishFindingRetestEmitsStatusChange(t *testing.T) {
d := notifyTestDB(t)
ctx := context.Background()
tk, err := d.CreateTask("复测推送测试", "目标", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
defer d.DeleteTask(tk.ID)
es := d.Exploration(tk.ExplorationID)
f, err := es.RecordFinding(ctx, RecordFindingInput{
TaskID: tk.ID, Worker: "test", VulnClass: "SQL注入", Name: "复测目标",
Severity: "high", Summary: "摘要",
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { d.Exec(`DELETE FROM notification_events WHERE finding_id=$1`, f.FindingID) })
// 建一条复测记录并直接推到完成态。
rt, _, _, err := d.CreateFindingRetest(ctx, f.FindingID, "复核")
if err != nil {
t.Fatal(err)
}
if rt.ConversationID == nil {
t.Fatal("复测应关联一个会话")
}
// 复测必须先进入 running 才能落结论(与真实流程一致)。
if ok, err := d.StartFindingRetest(ctx, rt.ID); err != nil || !ok {
t.Fatalf("启动复测失败: ok=%v err=%v", ok, err)
}
if err := d.RecordFindingRetestResult(ctx, *rt.ConversationID, "fixed", "已修复", "证据"); err != nil {
t.Fatal(err)
}
if err := d.FinishFindingRetest(rt.ID, "completed", ""); err != nil {
t.Fatalf("结束复测失败: %v", err)
}
var status string
if err := d.QueryRow(`SELECT status FROM findings WHERE id=$1`, f.FindingID).Scan(&status); err != nil {
t.Fatal(err)
}
if status != FindingFixed {
t.Fatalf("复测判已修复后状态应为 fixed,得到 %s", status)
}
// 关键断言:必须有一条状态变更事件,且 from/to 正确。
var snapshot []byte
err = d.QueryRow(`SELECT snapshot FROM notification_events WHERE finding_id=$1 AND kind=$2 ORDER BY id DESC LIMIT 1`,
f.FindingID, notify.EventFindingStatusChanged).Scan(&snapshot)
if err != nil {
t.Fatalf("复测判已修复应登记状态变更推送事件(否则配了 on_status_change 的渠道收不到): %v", err)
}
var snap notify.Snapshot
if err := json.Unmarshal(snapshot, &snap); err != nil {
t.Fatal(err)
}
if snap.FromStatus != "pending" || snap.ToStatus != FindingFixed {
t.Fatalf("快照的状态流转不对: %s → %s", snap.FromStatus, snap.ToStatus)
}
// 快照要带渲染所需字段,否则推送出来是空壳。
if snap.Name != "复测目标" || snap.Severity != "high" {
t.Fatalf("快照缺少渲染字段: %+v", snap)
}
}
+1386
View File
File diff suppressed because it is too large Load Diff
+43
View File
@@ -0,0 +1,43 @@
package db
import "database/sql"
// Settings is a tiny key-value store for global app config the UI toggles at
// runtime (e.g. traffic_capture). Missing keys fall back to caller defaults.
// GetSetting returns the stored value and ok=false when the key is unset.
func (d *DB) GetSetting(key string) (value string, ok bool, err error) {
err = d.QueryRow(`SELECT value FROM settings WHERE key=$1`, key).Scan(&value)
if err == sql.ErrNoRows {
return "", false, nil
}
if err != nil {
return "", false, err
}
return value, true, nil
}
// SetSetting upserts a setting value.
func (d *DB) SetSetting(key, value string) error {
_, err := d.Exec(`
INSERT INTO settings(key, value) VALUES ($1, $2)
ON CONFLICT (key) DO UPDATE SET value = EXCLUDED.value, updated_at = now()`, key, value)
return err
}
// GetBool returns the boolean setting, or def when unset/unparseable.
func (d *DB) GetBool(key string, def bool) bool {
v, ok, err := d.GetSetting(key)
if err != nil || !ok {
return def
}
return v == "true" || v == "1"
}
// SetBool stores a boolean setting as "true"/"false".
func (d *DB) SetBool(key string, val bool) error {
if val {
return d.SetSetting(key, "true")
}
return d.SetSetting(key, "false")
}
+288
View File
@@ -0,0 +1,288 @@
package db
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"github.com/Autumn-27/artex/sidequestion"
"github.com/google/uuid"
)
var ErrSideBusy = errors.New("当前会话已有旁路问题正在回答")
var ErrSideParentGone = errors.New("旁路父会话已删除或归档")
// Lock the real parent before the side session, also covering soft task/intent
// deletion. A delayed checkpoint cannot recreate data after archive cleanup.
func lockSideParent(ctx context.Context, tx *sql.Tx, p sidequestion.Parent) error {
var id int64
var err error
if p.ConversationID > 0 {
err = tx.QueryRowContext(ctx, `SELECT id FROM conversations WHERE id=$1 FOR SHARE`, p.ConversationID).Scan(&id)
} else {
err = tx.QueryRowContext(ctx, `SELECT id FROM tasks WHERE id=$1 AND exploration_id=$2 AND deleted_at IS NULL AND archived_at IS NULL FOR SHARE`, p.TaskID, p.ExplorationID).Scan(&id)
if err == nil && p.IntentID > 0 {
err = tx.QueryRowContext(ctx, `SELECT id FROM exploration_nodes WHERE id=$1 AND exploration_id=$2 AND kind='intent' AND state<>'stopped' FOR SHARE`, p.IntentID, p.ExplorationID).Scan(&id)
}
}
if errors.Is(err, sql.ErrNoRows) {
return ErrSideParentGone
}
return err
}
func nullableSideID(id int64) any {
if id == 0 {
return nil
}
return id
}
func (d *DB) SaveSideSnapshot(ctx context.Context, s sidequestion.Snapshot) error {
b, err := json.Marshal(s)
if err != nil {
return err
}
tx, err := d.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
if err = lockSideParent(ctx, tx, s.Parent); err != nil {
return err
}
_, err = tx.ExecContext(ctx, `INSERT INTO side_question_sessions(session_key,conversation_id,task_id,exploration_id,intent_id,run_id,version,snapshot)
VALUES($1,$2,$3,$4,$5,$6,$7,$8) ON CONFLICT(session_key) DO UPDATE SET run_id=EXCLUDED.run_id,version=EXCLUDED.version,snapshot=EXCLUDED.snapshot
WHERE (side_question_sessions.run_id,side_question_sessions.version)<(EXCLUDED.run_id,EXCLUDED.version)`,
s.Parent.Key(), nullableSideID(s.Parent.ConversationID), nullableSideID(s.Parent.TaskID), nullableSideID(s.Parent.ExplorationID), nullableSideID(s.Parent.IntentID), s.RunID, s.Version, string(jsonbClean(b)))
if err != nil {
return err
}
return tx.Commit()
}
func (d *DB) SideSnapshot(ctx context.Context, key string) (*sidequestion.Snapshot, error) {
var b []byte
err := d.QueryRowContext(ctx, `SELECT snapshot FROM side_question_sessions WHERE session_key=$1`, key).Scan(&b)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
if err != nil {
return nil, err
}
var s sidequestion.Snapshot
err = json.Unmarshal(b, &s)
return &s, err
}
const sideCols = `id,ordinal,session_key,generation,client_id,question,answer,status,error,model,snapshot_at,created_at,sequence,usage,context_info`
func (d *DB) ExistingSideRequest(ctx context.Context, key, client string) (*sidequestion.Exchange, error) {
e, err := scanSide(d.QueryRowContext(ctx, `SELECT `+sideCols+` FROM side_question_requests WHERE session_key=$1 AND client_id=$2 AND generation=(SELECT generation FROM side_question_sessions WHERE session_key=$1)`, key, client))
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
return &e, err
}
func scanSide(row interface{ Scan(...any) error }) (sidequestion.Exchange, error) {
var e sidequestion.Exchange
var model, usage, info []byte
err := row.Scan(&e.ID, &e.Ordinal, &e.SessionKey, &e.Generation, &e.ClientID, &e.Question, &e.Answer, &e.Status, &e.Error, &model, &e.SnapshotAt, &e.CreatedAt, &e.Sequence, &usage, &info)
if err != nil {
return e, err
}
if err = json.Unmarshal(model, &e.Model); err != nil {
return e, err
}
if err = json.Unmarshal(info, &e.Context); err != nil {
return e, err
}
err = json.Unmarshal(usage, &e.Usage)
return e, err
}
func (d *DB) SideRequest(ctx context.Context, id string) (*sidequestion.Exchange, error) {
e, err := scanSide(d.QueryRowContext(ctx, `SELECT `+sideCols+` FROM side_question_requests WHERE id=$1`, id))
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
return &e, err
}
func (d *DB) CurrentSideRequest(ctx context.Context, key string) (*sidequestion.Exchange, error) {
e, err := scanSide(d.QueryRowContext(ctx, `SELECT `+sideCols+` FROM side_question_requests WHERE session_key=$1 AND status='running'`, key))
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
return &e, err
}
func (d *DB) SideHistory(ctx context.Context, key string, before int64, limit int) ([]sidequestion.Exchange, error) {
rows, err := d.QueryContext(ctx, `SELECT `+sideCols+` FROM side_question_requests WHERE session_key=$1 AND ($2::bigint=0 OR ordinal<$2) ORDER BY ordinal DESC LIMIT $3`, key, before, limit)
if err != nil {
return nil, err
}
defer rows.Close()
out := []sidequestion.Exchange{}
for rows.Next() {
e, err := scanSide(rows)
if err != nil {
return nil, err
}
out = append(out, e)
}
return out, rows.Err()
}
func (d *DB) SideReplay(ctx context.Context, key string) ([]sidequestion.Exchange, error) {
rows, err := d.QueryContext(ctx, `SELECT `+sideCols+` FROM side_question_requests WHERE session_key=$1 AND status='completed' ORDER BY ordinal DESC LIMIT 20`, key)
if err != nil {
return nil, err
}
defer rows.Close()
out := []sidequestion.Exchange{}
for rows.Next() {
e, err := scanSide(rows)
if err != nil {
return nil, err
}
out = append(out, e)
}
for i, j := 0, len(out)-1; i < j; i, j = i+1, j-1 {
out[i], out[j] = out[j], out[i]
}
return out, rows.Err()
}
func (d *DB) StartSideRequest(ctx context.Context, s sidequestion.Snapshot, clientID, question string) (*sidequestion.Exchange, bool, error) {
tx, err := d.BeginTx(ctx, nil)
if err != nil {
return nil, false, err
}
defer tx.Rollback()
if err = lockSideParent(ctx, tx, s.Parent); err != nil {
return nil, false, err
}
var generation int64
if err = tx.QueryRowContext(ctx, `SELECT generation FROM side_question_sessions WHERE session_key=$1 FOR UPDATE`, s.Parent.Key()).Scan(&generation); err != nil {
return nil, false, err
}
e, err := scanSide(tx.QueryRowContext(ctx, `SELECT `+sideCols+` FROM side_question_requests WHERE session_key=$1 AND generation=$2 AND client_id=$3`, s.Parent.Key(), generation, clientID))
if err == nil {
if e.Question != question {
return nil, false, fmt.Errorf("同一请求 ID 不能用于不同问题")
}
return &e, false, tx.Commit()
}
if !errors.Is(err, sql.ErrNoRows) {
return nil, false, err
}
var busy bool
if err = tx.QueryRowContext(ctx, `SELECT EXISTS(SELECT 1 FROM side_question_requests WHERE session_key=$1 AND status='running')`, s.Parent.Key()).Scan(&busy); err != nil {
return nil, false, err
}
if busy {
return nil, false, ErrSideBusy
}
model, err := json.Marshal(s.Model)
if err != nil {
return nil, false, err
}
e, err = scanSide(tx.QueryRowContext(ctx, `INSERT INTO side_question_requests(id,session_key,generation,client_id,question,status,model,snapshot_at) VALUES($1,$2,$3,$4,$5,'running',$6,$7) RETURNING `+sideCols, uuid.NewString(), s.Parent.Key(), generation, clientID, question, string(model), s.CapturedAt))
if err != nil {
return nil, false, err
}
return &e, true, tx.Commit()
}
// Conditional updates cannot resurrect deleted history or overwrite a terminal
// cancellation with a late provider callback.
func (d *DB) UpdateSideRequest(ctx context.Context, e sidequestion.Exchange) (bool, error) {
usage, err := json.Marshal(e.Usage)
if err != nil {
return false, err
}
info, err := json.Marshal(e.Context)
if err != nil {
return false, err
}
r, err := d.ExecContext(ctx, `UPDATE side_question_requests r SET answer=$2,status=$3,error=$4,sequence=$5,usage=$6,context_info=$7
WHERE r.id=$1 AND r.status='running' AND r.sequence<$5 AND EXISTS(SELECT 1 FROM side_question_sessions s WHERE s.session_key=r.session_key AND s.generation=r.generation)`, e.ID, e.Answer, e.Status, e.Error, e.Sequence, string(usage), string(info))
if err != nil {
return false, err
}
n, err := r.RowsAffected()
return n == 1, err
}
func (d *DB) ClearSideHistory(ctx context.Context, key string) error {
tx, err := d.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
if _, err = tx.ExecContext(ctx, `UPDATE side_question_sessions SET generation=generation+1,memory='{}' WHERE session_key=$1`, key); err != nil {
return err
}
if _, err = tx.ExecContext(ctx, `DELETE FROM side_question_requests WHERE session_key=$1`, key); err != nil {
return err
}
return tx.Commit()
}
func (d *DB) SideMemory(ctx context.Context, e sidequestion.Exchange) (sidequestion.Memory, error) {
var memory sidequestion.Memory
var raw []byte
err := d.QueryRowContext(ctx, `SELECT s.memory FROM side_question_sessions s JOIN side_question_requests r ON r.session_key=s.session_key
WHERE r.id=$1 AND r.generation=s.generation AND s.generation=$2 AND r.status='running'`, e.ID, e.Generation).Scan(&raw)
if err != nil {
return memory, err
}
err = json.Unmarshal(raw, &memory)
return memory, err
}
// Unlike SideReplay's UI-era 20-row window, this cursor visits all unsummarized
// successful exchanges, in bounded pages and only before the admitted request.
func (d *DB) SideReplayPage(ctx context.Context, e sidequestion.Exchange, after int64) ([]sidequestion.Exchange, error) {
rows, err := d.QueryContext(ctx, `SELECT `+sideCols+` FROM side_question_requests WHERE session_key=$1 AND generation=$2
AND status='completed' AND ordinal>$3 AND ordinal<$4 ORDER BY ordinal LIMIT 20`, e.SessionKey, e.Generation, after, e.Ordinal)
if err != nil {
return nil, err
}
defer rows.Close()
var out []sidequestion.Exchange
for rows.Next() {
item, err := scanSide(rows)
if err != nil {
return nil, err
}
out = append(out, item)
}
return out, rows.Err()
}
func (d *DB) SaveSideMemory(ctx context.Context, e sidequestion.Exchange, memory sidequestion.Memory) error {
raw, err := json.Marshal(memory)
if err != nil {
return err
}
result, err := d.ExecContext(ctx, `UPDATE side_question_sessions s SET memory=$3 WHERE s.session_key=$1 AND s.generation=$2
AND EXISTS(SELECT 1 FROM side_question_requests r WHERE r.id=$4 AND r.session_key=s.session_key AND r.generation=s.generation AND r.status='running')`, e.SessionKey, e.Generation, string(raw), e.ID)
if err != nil {
return err
}
n, err := result.RowsAffected()
if err == nil && n == 0 {
return ErrSideParentGone
}
return err
}
func (d *DB) InterruptSideRequests(ctx context.Context) error {
_, err := d.ExecContext(ctx, `UPDATE side_question_requests SET status='interrupted',error='服务重启,回答已中断',sequence=sequence+1 WHERE status='running'`)
return err
}
+446
View File
@@ -0,0 +1,446 @@
package db
import (
"context"
"encoding/json"
"errors"
"fmt"
"sync"
"testing"
"time"
"github.com/Autumn-27/artex/sidequestion"
"github.com/Autumn-27/norma/llm"
)
func sideFixture(t *testing.T) (*DB, sidequestion.Snapshot) {
t.Helper()
d, err := Open(testDSN(t))
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { d.Close() })
c, err := d.CreateConversation("mainagent", "side persistence", nil)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteConversation(c.ID) })
s := sidequestion.Snapshot{Parent: sidequestion.Parent{ConversationID: c.ID}, RunID: 1, Version: 1, CapturedAt: time.Now().UTC(), Model: sidequestion.Model{Model: "fixture"}, Request: llm.CompletionRequest{Messages: []llm.Message{llm.UserText("main-only")}}}
if err = d.SaveSideSnapshot(t.Context(), s); err != nil {
t.Fatal(err)
}
return d, s
}
func TestSideHistoryIdempotencyPagingAndRecovery(t *testing.T) {
d, s := sideFixture(t)
ctx := t.Context()
first, created, err := d.StartSideRequest(ctx, s, "request-0", "question-0")
if err != nil || !created {
t.Fatalf("start %v %v", created, err)
}
again, created, err := d.StartSideRequest(ctx, s, "request-0", "question-0")
if err != nil || created || again.ID != first.ID {
t.Fatalf("dedup %v %v", created, err)
}
if _, _, err = d.StartSideRequest(ctx, s, "request-0", "different"); err == nil {
t.Fatal("conflicting duplicate accepted")
}
if _, _, err = d.StartSideRequest(ctx, s, "request-1", "question-1"); !errors.Is(err, ErrSideBusy) {
t.Fatalf("busy %v", err)
}
for i := 0; i < 24; i++ {
e := first
if i > 0 {
e, _, err = d.StartSideRequest(ctx, s, fmt.Sprintf("request-%d", i), fmt.Sprintf("question-%d", i))
if err != nil {
t.Fatal(err)
}
}
e.Answer = fmt.Sprintf("answer-%d", i)
e.Sequence = 1
e.Status = "completed"
if ok, err := d.UpdateSideRequest(ctx, *e); err != nil || !ok {
t.Fatalf("finish %v %v", ok, err)
}
}
page, err := d.SideHistory(ctx, s.Parent.Key(), 0, 20)
if err != nil || len(page) != 20 {
t.Fatalf("page %d %v", len(page), err)
}
tail, err := d.SideHistory(ctx, s.Parent.Key(), page[19].Ordinal, 20)
if err != nil || len(tail) != 4 || tail[0].Ordinal >= page[19].Ordinal {
t.Fatalf("tail %+v %v", tail, err)
}
replay, err := d.SideReplay(ctx, s.Parent.Key())
if err != nil || len(replay) != 20 || replay[0].Question != "question-4" || replay[19].Question != "question-23" {
t.Fatalf("replay %+v %v", replay, err)
}
e, _, err := d.StartSideRequest(ctx, s, "unfinished", "partial question")
if err != nil {
t.Fatal(err)
}
e.Answer = "saved partial"
e.Sequence = 1
e.Usage.InputTokens = 17
if _, err = d.UpdateSideRequest(ctx, *e); err != nil {
t.Fatal(err)
}
if err = d.InterruptSideRequests(ctx); err != nil {
t.Fatal(err)
}
got, err := d.SideRequest(ctx, e.ID)
if err != nil || got.Status != "interrupted" || got.Answer != "saved partial" || got.Usage.InputTokens != 17 {
t.Fatalf("recovery %+v %v", got, err)
}
saved, err := d.SideSnapshot(ctx, s.Parent.Key())
if err != nil || saved.Request.Messages[0].Text() != "main-only" {
t.Fatalf("snapshot %+v %v", saved, err)
}
if _, _, err = d.StartSideRequest(ctx, *saved, "after-restart", "continue"); err != nil {
t.Fatal(err)
}
}
func TestSideMemoryPagingClearAndRestart(t *testing.T) {
d, s := sideFixture(t)
ctx := t.Context()
var ordinal int64
for i := 0; i < 50; i++ {
e, _, err := d.StartSideRequest(ctx, s, fmt.Sprint(i), "history")
if err != nil {
t.Fatal(err)
}
e.Status = "completed"
e.Sequence = 1
e.Answer = "saved"
if i == 2 {
e.Status = "failed"
}
if _, err = d.UpdateSideRequest(ctx, *e); err != nil {
t.Fatal(err)
}
if i == 29 {
ordinal = e.Ordinal
}
}
e, _, err := d.StartSideRequest(ctx, s, "admitted", "question")
if err != nil {
t.Fatal(err)
}
memory := sidequestion.Memory{History: "old decision", Through: ordinal, SnapshotKey: "snapshot", SnapshotSummary: "main evidence", TailStart: 3}
if err = d.SaveSideMemory(ctx, *e, memory); err != nil {
t.Fatal(err)
}
var all []sidequestion.Exchange
for after := int64(0); ; {
page, err := d.SideReplayPage(ctx, *e, after)
if err != nil {
t.Fatal(err)
}
if len(page) == 0 {
break
}
if len(page) > 20 {
t.Fatal("unbounded page")
}
all = append(all, page...)
after = page[len(page)-1].Ordinal
}
if len(all) != 49 {
t.Fatalf("history missing/duplicated: %d", len(all))
}
if err = d.InterruptSideRequests(ctx); err != nil {
t.Fatal(err)
}
next, _, err := d.StartSideRequest(ctx, s, "restart", "continue")
if err != nil {
t.Fatal(err)
}
got, err := d.SideMemory(ctx, *next)
if err != nil || got != memory {
t.Fatalf("memory after restart: %+v %v", got, err)
}
page, err := d.SideReplayPage(ctx, *next, memory.Through)
if err != nil || len(page) != 20 || page[0].Ordinal <= ordinal {
t.Fatalf("summary cursor: %+v %v", page, err)
}
if err = d.ClearSideHistory(ctx, s.Parent.Key()); err != nil {
t.Fatal(err)
}
if err = d.SaveSideMemory(ctx, *next, memory); !errors.Is(err, ErrSideParentGone) {
t.Fatalf("late memory resurrected: %v", err)
}
fresh, _, err := d.StartSideRequest(ctx, s, "after-clear", "fresh")
if err != nil {
t.Fatal(err)
}
got, err = d.SideMemory(ctx, *fresh)
if err != nil || got != (sidequestion.Memory{}) {
t.Fatalf("memory survived clear: %+v %v", got, err)
}
if snapshot, err := d.SideSnapshot(ctx, s.Parent.Key()); err != nil || snapshot.Request.Messages[0].Text() != "main-only" {
t.Fatal("memory changed main snapshot")
}
}
func TestSideMemoryClearRace(t *testing.T) {
d, s := sideFixture(t)
for i := 0; i < 10; i++ {
e, _, err := d.StartSideRequest(t.Context(), s, "race", "question")
if err != nil {
t.Fatal(err)
}
var wg sync.WaitGroup
wg.Add(2)
go func() {
defer wg.Done()
if err := d.ClearSideHistory(t.Context(), s.Parent.Key()); err != nil {
t.Error(err)
}
}()
go func() {
defer wg.Done()
err := d.SaveSideMemory(t.Context(), *e, sidequestion.Memory{History: "late", Through: e.Ordinal - 1})
if err != nil && !errors.Is(err, ErrSideParentGone) {
t.Error(err)
}
}()
wg.Wait()
var raw []byte
if err = d.QueryRow(`SELECT memory FROM side_question_sessions WHERE session_key=$1`, s.Parent.Key()).Scan(&raw); err != nil || string(raw) != "{}" {
t.Fatalf("late cache write: %s %v", raw, err)
}
}
}
func TestSideArchiveRowsWithoutNewFields(t *testing.T) {
d, s := sideFixture(t)
e, _, err := d.StartSideRequest(t.Context(), s, "legacy", "legacy question")
if err != nil {
t.Fatal(err)
}
tx, err := d.Begin()
if err != nil {
t.Fatal(err)
}
defer tx.Rollback()
rows := map[string]json.RawMessage{}
for _, table := range []string{"side_question_sessions", "side_question_requests"} {
var raw []byte
if err = tx.QueryRow(`SELECT json_agg(t) FROM `+table+` t WHERE session_key=$1`, s.Parent.Key()).Scan(&raw); err != nil {
t.Fatal(err)
}
var items []map[string]json.RawMessage
if err = json.Unmarshal(raw, &items); err != nil {
t.Fatal(err)
}
for _, item := range items {
delete(item, "memory")
delete(item, "context_info")
}
rows[table], err = json.Marshal(items)
if err != nil {
t.Fatal(err)
}
}
if _, err = tx.Exec(`DELETE FROM side_question_sessions WHERE session_key=$1`, s.Parent.Key()); err != nil {
t.Fatal(err)
}
for _, table := range []string{"side_question_sessions", "side_question_requests"} {
if err = insertArchiveRows(tx, table, rows[table]); err != nil {
t.Fatal(err)
}
}
var info, memory []byte
if err = tx.QueryRow(`SELECT context_info,memory FROM side_question_requests r JOIN side_question_sessions s USING(session_key) WHERE r.id=$1`, e.ID).Scan(&info, &memory); err != nil || string(info) != "{}" || string(memory) != "{}" {
t.Fatalf("legacy defaults: %s %s %v", info, memory, err)
}
}
func TestSideClearLateWritersAndDeletedParent(t *testing.T) {
d, s := sideFixture(t)
ctx := t.Context()
for i := 0; i < 10; i++ {
e, _, err := d.StartSideRequest(ctx, s, "same-client", "question")
if err != nil {
t.Fatal(err)
}
var wg sync.WaitGroup
wg.Add(2)
go func() {
defer wg.Done()
if err := d.ClearSideHistory(ctx, s.Parent.Key()); err != nil {
t.Error(err)
}
}()
go func() {
defer wg.Done()
copy := *e
copy.Sequence = 1
copy.Answer = "late"
copy.Status = "completed"
if _, err := d.UpdateSideRequest(ctx, copy); err != nil {
t.Error(err)
}
}()
wg.Wait()
if row, err := d.SideRequest(ctx, e.ID); err != nil || row != nil {
t.Fatalf("cleared answer resurrected: %+v %v", row, err)
}
e.Sequence = 2
e.Status = "completed"
if ok, err := d.UpdateSideRequest(ctx, *e); err != nil || ok {
t.Fatalf("late update %v %v", ok, err)
}
}
s.Version = 3
if err := d.SaveSideSnapshot(ctx, s); err != nil {
t.Fatal(err)
}
s.Version = 2
if err := d.SaveSideSnapshot(ctx, s); err != nil {
t.Fatal(err)
}
got, err := d.SideSnapshot(ctx, s.Parent.Key())
if err != nil || got.Version != 3 {
t.Fatalf("older version won: %+v %v", got, err)
}
if err = d.DeleteConversation(s.Parent.ConversationID); err != nil {
t.Fatal(err)
}
if err = d.SaveSideSnapshot(ctx, s); !errors.Is(err, ErrSideParentGone) {
t.Fatalf("deleted parent restored: %v", err)
}
if got, err = d.SideSnapshot(ctx, s.Parent.Key()); err != nil || got != nil {
t.Fatalf("delete cascade: %+v %v", got, err)
}
}
func TestSideTaskArchiveVersions(t *testing.T) {
for _, version := range []int{1, 2, 3} {
t.Run(fmt.Sprint(version), func(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Fatal(err)
}
defer d.Close()
if err := d.EnsureLLMRecordsTable(); err != nil {
t.Fatal(err)
}
if err := d.EnsureLLMUsageTable(); err != nil {
t.Fatal(err)
}
task, err := d.CreateTask("btw archive", "restore context", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
defer func() {
_, _ = d.Exec(`DELETE FROM task_archives WHERE task_id=$1`, task.ID)
_ = d.DeleteTask(task.ID)
}()
iid, err := d.Exploration(task.ExplorationID).AddNode(KindIntent, map[string]any{"summary": "worker"}, 1, "paused", "planner", nil)
if err != nil {
t.Fatal(err)
}
var snapshots []sidequestion.Snapshot
memories := make(map[string]sidequestion.Memory)
contextInfo := sidequestion.ContextInfo{Phase: "answering", RecentExchanges: 20, HistorySummarized: true, SnapshotSummarized: true, EstimatedInputTokens: 12000, InputBudget: 16000, OutputTokens: 2048}
for _, intent := range []int64{0, iid} {
s := sidequestion.Snapshot{Parent: sidequestion.Parent{TaskID: task.ID, ExplorationID: task.ExplorationID, IntentID: intent}, RunID: 1, Version: 2, CapturedAt: time.Now().UTC(), Request: llm.CompletionRequest{Messages: []llm.Message{llm.UserText("archived main context")}}}
if err = d.SaveSideSnapshot(t.Context(), s); err != nil {
t.Fatal(err)
}
e, _, err := d.StartSideRequest(t.Context(), s, "client", "archive question")
if err != nil {
t.Fatal(err)
}
memory := sidequestion.Memory{History: "archived early decision", Through: e.Ordinal, SnapshotKey: s.Parent.Key(), SnapshotSummary: "archived evidence", TailStart: 1}
if err = d.SaveSideMemory(t.Context(), *e, memory); err != nil {
t.Fatal(err)
}
memories[s.Parent.Key()] = memory
e.Answer = "archive answer"
e.Sequence = 1
e.Status = "completed"
e.Context = contextInfo
if _, err = d.UpdateSideRequest(t.Context(), *e); err != nil {
t.Fatal(err)
}
snapshots = append(snapshots, s)
}
if err := d.SetPaused(task.ID, true); err != nil {
t.Fatal(err)
}
job, err := d.QueueTaskArchive(task.ID)
if err != nil {
t.Fatal(err)
}
if _, err = d.ClaimTaskArchiveJob(t.Context()); err != nil {
t.Fatal(err)
}
archive, err := d.SnapshotTaskArchive(task.ID)
if err != nil {
t.Fatal(err)
}
if archive.FormatVersion != 3 || archive.DataCounts["side_question_sessions"] != 2 || archive.DataCounts["side_question_requests"] != 2 {
t.Fatalf("missing side archive: %+v", archive.DataCounts)
}
if err = d.CompleteTaskArchive(job.ID, archive, "/tmp/side-fixture.tar.zst", "fixture", 1, 1); err != nil {
t.Fatal(err)
}
if err = d.SaveSideSnapshot(t.Context(), snapshots[0]); !errors.Is(err, ErrSideParentGone) {
t.Fatalf("late archived snapshot: %v", err)
}
if got, err := d.SideSnapshot(t.Context(), snapshots[0].Parent.Key()); err != nil || got != nil {
t.Fatalf("archive retained hot snapshot %+v %v", got, err)
}
if _, err = d.QueueTaskArchiveRestore(job.ID); err != nil {
t.Fatal(err)
}
if _, err = d.ClaimTaskArchiveJob(t.Context()); err != nil {
t.Fatal(err)
}
archive.FormatVersion = version
if version < 3 {
delete(archive.Tables, "side_question_sessions")
delete(archive.Tables, "side_question_requests")
delete(archive.DataCounts, "side_question_sessions")
delete(archive.DataCounts, "side_question_requests")
}
if _, err = d.RestoreTaskArchive(job.ID, archive, 0); err != nil {
t.Fatal(err)
}
for _, s := range snapshots {
got, err := d.SideSnapshot(context.Background(), s.Parent.Key())
if err != nil {
t.Fatal(err)
}
if version < 3 {
if got != nil {
t.Fatal("legacy archive fabricated snapshot")
}
} else {
if got == nil || got.Request.Messages[0].Text() != "archived main context" {
t.Fatalf("restored snapshot: %+v", got)
}
history, err := d.SideHistory(t.Context(), s.Parent.Key(), 0, 20)
if err != nil || len(history) != 1 || history[0].Answer != "archive answer" {
t.Fatalf("restored history %+v %v", history, err)
}
if history[0].Context != contextInfo {
t.Fatalf("restored context metadata: %+v", history[0].Context)
}
next, _, err := d.StartSideRequest(t.Context(), *got, "after-restore", "continue")
if err != nil {
t.Fatal(err)
}
memory, err := d.SideMemory(t.Context(), *next)
if err != nil || memory != memories[s.Parent.Key()] {
t.Fatalf("restored summary cache: %+v %v", memory, err)
}
}
}
})
}
}
+217
View File
@@ -0,0 +1,217 @@
package db
import (
"database/sql"
"strings"
"time"
)
// SkillUsage is one Skill() invocation — the always-on skill call ledger. Written
// from the Skill meta-tool's OnInvoke hook (server/assembly.go), one row per load.
// Carries only dimensions (which skill, which agent, which task/session), never the
// caller's args text — args_len is kept so an "empty vs substantial context" split
// is still possible without storing prompt content, mirroring llm_usage.
//
// Rows deliberately outlive their task: skill_usage has no foreign keys, so deleting
// a task keeps its skill statistics intact (same rationale as llm_usage).
type SkillUsage struct {
Skill string `json:"skill"` // skill directory name (matches agent_skill_visibility.skill_name)
AgentKey string `json:"agent_key"` // worker / planner / mainagent / custom agent key
TaskID int64 `json:"task_id"` // 0 for non-task runs (chat sessions)
ExplorationID int64 `json:"exploration_id"` // 0 when unknown
IntentID int64 `json:"intent_id"` // worker's intent node; 0 for planner/mainagent/chat
SessionID string `json:"session_id"` // chat conversation id; empty for task runs
ArgsLen int `json:"args_len"`
Found bool `json:"found"` // false = the model named a skill that does not exist
}
// InsertSkillUsage appends one ledger row. Best-effort: callers log and continue on
// error (a lost metering row must never break a skill invocation).
func (d *DB) InsertSkillUsage(u *SkillUsage) error {
_, err := d.Exec(`
INSERT INTO skill_usage(skill, agent_key, task_id, exploration_id, intent_id, session_id, args_len, found)
VALUES ($1, NULLIF($2,''), $3, $4, $5, NULLIF($6,''), $7, $8)`,
u.Skill, u.AgentKey, nullIfZero(u.TaskID), nullIfZero(u.ExplorationID),
nullIfZero(u.IntentID), u.SessionID, u.ArgsLen, u.Found)
return err
}
func nullIfZero(v int64) any {
if v > 0 {
return v
}
return nil
}
// SkillStat is one skill's aggregate usage, for the skills page.
type SkillStat struct {
Skill string `json:"skill"`
Calls int `json:"calls"`
Tasks int `json:"tasks"` // distinct tasks that loaded it (chat runs excluded)
Agents []string `json:"agents"` // agent keys that loaded it, most-used first
LastUsed *time.Time `json:"last_used"` // nil when never called
}
// SkillStats aggregates the whole ledger grouped by skill, most-used first. Skills
// that were never invoked are absent — callers merge against the skill list on disk.
// Only resolved calls count; misses are reported separately by MissingSkillStats.
func (d *DB) SkillStats() ([]SkillStat, error) {
// agent keys come back as one comma-joined string rather than text[]: the pgx
// stdlib driver has no database/sql Scan target for arrays, and agent keys are
// [a-z0-9_-] so a comma join is unambiguous.
rows, err := d.Query(`
SELECT skill, COUNT(*) AS calls,
COUNT(DISTINCT task_id) AS tasks,
COALESCE(STRING_AGG(DISTINCT agent_key, ','), '') AS agents,
MAX(ts) AS last_used
FROM skill_usage
WHERE found
GROUP BY skill
ORDER BY COUNT(*) DESC, skill`)
if err != nil {
return nil, err
}
defer rows.Close()
out := []SkillStat{}
for rows.Next() {
var s SkillStat
var agents string
var lastUsed sql.NullTime
if err := rows.Scan(&s.Skill, &s.Calls, &s.Tasks, &agents, &lastUsed); err != nil {
return nil, err
}
s.Agents = []string{}
if agents != "" {
s.Agents = strings.Split(agents, ",")
}
if lastUsed.Valid {
t := lastUsed.Time
s.LastUsed = &t
}
out = append(out, s)
}
if err := rows.Err(); err != nil {
return nil, err
}
archived, err := d.archivedTaskAggregates()
if err != nil {
return nil, err
}
return mergeArchivedSkillStats(out, archived, false), nil
}
// MissingSkillStats returns the skill names agents asked for that do not exist,
// most-requested first — the "wished it existed" gap list. Names come from the model
// so they are shown as-is (already length-capped at insert time).
func (d *DB) MissingSkillStats(limit int) ([]SkillStat, error) {
if limit <= 0 || limit > 200 {
limit = 20
}
rows, err := d.Query(`
SELECT skill, COUNT(*) AS calls,
COALESCE(STRING_AGG(DISTINCT agent_key, ','), '') AS agents,
MAX(ts) AS last_used
FROM skill_usage
WHERE NOT found
GROUP BY skill
ORDER BY COUNT(*) DESC, skill`)
if err != nil {
return nil, err
}
defer rows.Close()
out := []SkillStat{}
for rows.Next() {
var s SkillStat
var agents string
var lastUsed sql.NullTime
if err := rows.Scan(&s.Skill, &s.Calls, &agents, &lastUsed); err != nil {
return nil, err
}
s.Agents = []string{}
if agents != "" {
s.Agents = strings.Split(agents, ",")
}
if lastUsed.Valid {
t := lastUsed.Time
s.LastUsed = &t
}
out = append(out, s)
}
if err := rows.Err(); err != nil {
return nil, err
}
archived, err := d.archivedTaskAggregates()
if err != nil {
return nil, err
}
out = mergeArchivedSkillStats(out, archived, true)
if len(out) > limit {
out = out[:limit]
}
return out, nil
}
// SkillCall is one row of a skill's recent-call list (detail panel).
type SkillCall struct {
TS time.Time `json:"ts"`
AgentKey string `json:"agent_key"`
TaskID int64 `json:"task_id"`
SessionID string `json:"session_id"`
ArgsLen int `json:"args_len"`
}
// RecentSkillCalls returns the most recent invocations of one skill, newest first.
func (d *DB) RecentSkillCalls(skill string, limit int) ([]SkillCall, error) {
if limit <= 0 || limit > 200 {
limit = 50
}
rows, err := d.Query(`
SELECT ts, COALESCE(agent_key,''), COALESCE(task_id,0), COALESCE(session_id,''), args_len
FROM skill_usage
WHERE skill = $1 AND found
ORDER BY ts DESC
LIMIT $2`, skill, limit)
if err != nil {
return nil, err
}
defer rows.Close()
out := []SkillCall{}
for rows.Next() {
var c SkillCall
if err := rows.Scan(&c.TS, &c.AgentKey, &c.TaskID, &c.SessionID, &c.ArgsLen); err != nil {
return nil, err
}
out = append(out, c)
}
return out, rows.Err()
}
// SkillCallsByTask counts a task's skill loads, most-used first. Powers a per-task
// view of which procedures its agents actually reached for.
func (d *DB) SkillCallsByTask(taskID int64) ([]SkillStat, error) {
rows, err := d.Query(`
SELECT skill, COUNT(*) AS calls, MAX(ts) AS last_used
FROM skill_usage
WHERE task_id = $1 AND found
GROUP BY skill
ORDER BY COUNT(*) DESC, skill`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
out := []SkillStat{}
for rows.Next() {
var s SkillStat
var lastUsed sql.NullTime
if err := rows.Scan(&s.Skill, &s.Calls, &lastUsed); err != nil {
return nil, err
}
s.Agents = []string{}
if lastUsed.Valid {
t := lastUsed.Time
s.LastUsed = &t
}
out = append(out, s)
}
return out, rows.Err()
}
+116
View File
@@ -0,0 +1,116 @@
package db
import (
"testing"
)
// TestSkillUsageLedger writes a few rows and checks the three aggregates the
// skills page relies on: per-skill totals (resolved calls only), the miss list,
// and the recent-call list. Rows are namespaced by a unique skill name so the
// shared dev DB stays usable and the test cleans up after itself.
func TestSkillUsageLedger(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
const (
skillA = "zz-test-skill-a"
missing = "zz-test-skill-missing"
)
cleanup := func() {
_, _ = d.Exec(`DELETE FROM skill_usage WHERE skill IN ($1,$2)`, skillA, missing)
}
cleanup()
defer cleanup()
rows := []*SkillUsage{
{Skill: skillA, AgentKey: "worker", TaskID: 991, ExplorationID: 5, IntentID: 7, ArgsLen: 12, Found: true},
{Skill: skillA, AgentKey: "worker", TaskID: 991, ExplorationID: 5, ArgsLen: 0, Found: true},
{Skill: skillA, AgentKey: "planner", TaskID: 992, ExplorationID: 6, Found: true},
{Skill: skillA, AgentKey: "chatbot", SessionID: "conv-1", Found: true},
{Skill: missing, AgentKey: "worker", TaskID: 991, Found: false},
{Skill: missing, AgentKey: "worker", TaskID: 991, Found: false},
}
for _, r := range rows {
if err := d.InsertSkillUsage(r); err != nil {
t.Fatalf("insert %s: %v", r.Skill, err)
}
}
stats, err := d.SkillStats()
if err != nil {
t.Fatal(err)
}
var got *SkillStat
for i := range stats {
if stats[i].Skill == skillA {
got = &stats[i]
}
if stats[i].Skill == missing {
t.Fatalf("misses must not appear in SkillStats: %+v", stats[i])
}
}
if got == nil {
t.Fatalf("skill %s absent from SkillStats", skillA)
}
if got.Calls != 4 {
t.Errorf("calls: want 4, got %d", got.Calls)
}
// 991 + 992; the chat row has a NULL task_id and COUNT(DISTINCT) skips NULLs.
if got.Tasks != 2 {
t.Errorf("tasks: want 2, got %d", got.Tasks)
}
if len(got.Agents) != 3 {
t.Errorf("agents: want 3 distinct, got %v", got.Agents)
}
if got.LastUsed == nil {
t.Error("last_used must be set")
}
miss, err := d.MissingSkillStats(10)
if err != nil {
t.Fatal(err)
}
var missCalls int
for _, m := range miss {
if m.Skill == missing {
missCalls = m.Calls
}
}
if missCalls != 2 {
t.Errorf("missing calls: want 2, got %d", missCalls)
}
calls, err := d.RecentSkillCalls(skillA, 10)
if err != nil {
t.Fatal(err)
}
if len(calls) != 4 {
t.Fatalf("recent calls: want 4, got %d", len(calls))
}
// newest first
for i := 1; i < len(calls); i++ {
if calls[i].TS.After(calls[i-1].TS) {
t.Errorf("recent calls not newest-first at %d", i)
}
}
byTask, err := d.SkillCallsByTask(991)
if err != nil {
t.Fatal(err)
}
var taskCalls int
for _, s := range byTask {
if s.Skill == skillA {
taskCalls = s.Calls
}
if s.Skill == missing {
t.Errorf("misses must not appear in SkillCallsByTask: %+v", s)
}
}
if taskCalls != 2 {
t.Errorf("task 991 calls: want 2, got %d", taskCalls)
}
}
+91
View File
@@ -0,0 +1,91 @@
package db
import (
"encoding/json"
"sort"
)
// taskArchiveAggregate is the compact, non-sensitive statistics section of a
// cold archive manifest. Fields are additive so older archive formats remain
// readable when new dashboard dimensions are introduced.
type taskArchiveAggregate struct {
TokenProfiles []ProfileUsage `json:"token_profiles"`
TokenDaily []ProfileDayUsage `json:"token_daily"`
Skills map[string]int `json:"skills"`
SkillStats []SkillStat `json:"skill_stats"`
MissingSkillStats []SkillStat `json:"missing_skill_stats"`
Tools map[string]int `json:"tools"`
FindingStats FindingStats `json:"finding_stats"`
}
func (d *DB) archivedTaskAggregates() ([]taskArchiveAggregate, error) {
rawItems, err := d.ArchivedAggregateStats()
if err != nil {
return nil, err
}
out := make([]taskArchiveAggregate, 0, len(rawItems))
for _, raw := range rawItems {
var item taskArchiveAggregate
if err := json.Unmarshal(raw, &item); err != nil {
return nil, err
}
out = append(out, item)
}
return out, nil
}
func mergeAgentKeys(current, incoming []string) []string {
seen := make(map[string]bool, len(current)+len(incoming))
out := make([]string, 0, len(current)+len(incoming))
for _, group := range [][]string{current, incoming} {
for _, key := range group {
if key == "" || seen[key] {
continue
}
seen[key] = true
out = append(out, key)
}
}
return out
}
func mergeArchivedSkillStats(live []SkillStat, archived []taskArchiveAggregate, missing bool) []SkillStat {
byName := make(map[string]SkillStat, len(live))
for _, current := range live {
byName[current.Skill] = current
}
for _, aggregate := range archived {
coldStats := aggregate.SkillStats
if missing {
coldStats = aggregate.MissingSkillStats
} else if len(coldStats) == 0 {
// Format v1 initially retained only the call-count map. Preserve those
// summaries even though agent and timestamp dimensions are unavailable.
for skill, calls := range aggregate.Skills {
coldStats = append(coldStats, SkillStat{Skill: skill, Calls: calls, Tasks: 1})
}
}
for _, cold := range coldStats {
current := byName[cold.Skill]
current.Skill = cold.Skill
current.Calls += cold.Calls
current.Tasks += cold.Tasks
current.Agents = mergeAgentKeys(current.Agents, cold.Agents)
if cold.LastUsed != nil && (current.LastUsed == nil || cold.LastUsed.After(*current.LastUsed)) {
current.LastUsed = cold.LastUsed
}
byName[cold.Skill] = current
}
}
out := make([]SkillStat, 0, len(byName))
for _, current := range byName {
out = append(out, current)
}
sort.Slice(out, func(i, j int) bool {
if out[i].Calls != out[j].Calls {
return out[i].Calls > out[j].Calls
}
return out[i].Skill < out[j].Skill
})
return out
}
+787
View File
@@ -0,0 +1,787 @@
package db
import (
"bytes"
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"io"
"net/url"
"sort"
"strconv"
"strings"
"time"
)
const (
TaskArchiveFormatVersion = 3
TaskArchiveLegacyFormatVersion = 1
TaskArchiveLLMRecordsPath = "database/llm_records.ndjson"
)
func IsTaskArchiveFormatSupported(version int) bool {
return version >= TaskArchiveLegacyFormatVersion && version <= TaskArchiveFormatVersion
}
const (
ArchiveQueued = "archive_queued"
Archiving = "archiving"
ArchiveFailed = "archive_failed"
ArchiveReady = "ready"
RestoreQueued = "restore_queued"
Restoring = "restoring"
RestoreFailed = "restore_failed"
DeleteQueued = "delete_queued"
Deleting = "deleting"
DeleteFailed = "delete_failed"
)
var (
ErrTaskArchiveNotFound = errors.New("task archive not found")
ErrTaskArchiveIneligible = errors.New("task must be paused or terminal before archiving")
ErrTaskArchiveQueued = errors.New("queued task must be paused before archiving")
ErrTaskArchiveDependent = errors.New("task is inherited by a live task")
ErrTaskArchiveState = errors.New("task archive state does not allow this operation")
ErrTaskArchiveDeleteBlocked = errors.New("task archive is required by another archive")
ErrTaskArchiveFormatMismatch = errors.New("task archive format is not supported")
)
// TaskArchive is the compact PostgreSQL record retained while a task is cold.
// Sensitive profile configuration and API keys are intentionally absent.
type TaskArchive struct {
ID int64 `json:"id"`
TaskID int64 `json:"task_id"`
State string `json:"state"`
Phase string `json:"phase"`
Progress int `json:"progress"`
Error string `json:"error,omitempty"`
Warnings json.RawMessage `json:"warnings"`
FormatVersion int `json:"format_version"`
ArchivePath string `json:"-"`
SHA256 string `json:"sha256,omitempty"`
OriginalSize int64 `json:"original_size"`
CompressedSize int64 `json:"compressed_size"`
TaskName string `json:"task_name"`
TaskDescription string `json:"task_description"`
TaskGoal string `json:"task_goal"`
OriginalStatus string `json:"original_status"`
CategoryIDSnapshot *int64 `json:"category_id,omitempty"`
CategoryNameSnapshot string `json:"category_name,omitempty"`
SourceTaskIDs []int64 `json:"source_task_ids"`
RemainingTimeoutSeconds int64 `json:"remaining_timeout_seconds"`
DataCounts json.RawMessage `json:"data_counts"`
AggregateStats json.RawMessage `json:"aggregate_stats"`
ArchivedAt *time.Time `json:"archived_at,omitempty"`
RequestedAt time.Time `json:"requested_at"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// TaskArchiveBlockers returns one live direct dependent for every source task
// that cannot currently be archived. Dependents already queued for archiving do
// not block their source because the FIFO worker will compact them first.
func (d *DB) TaskArchiveBlockers() (map[int64]int64, error) {
rows, err := d.Query(`SELECT relation.source_task_id, MIN(child.id)
FROM task_relations relation
JOIN tasks child ON child.id=relation.task_id AND child.deleted_at IS NULL
LEFT JOIN task_archives pending ON pending.task_id=child.id
WHERE pending.id IS NULL OR pending.state NOT IN ('archive_queued','archiving')
GROUP BY relation.source_task_id`)
if err != nil {
return nil, err
}
defer rows.Close()
blockers := map[int64]int64{}
for rows.Next() {
var sourceID, dependentID int64
if err := rows.Scan(&sourceID, &dependentID); err != nil {
return nil, err
}
blockers[sourceID] = dependentID
}
return blockers, rows.Err()
}
type TaskArchivePage struct {
Items []TaskArchive `json:"items"`
Total int `json:"total"`
Page int `json:"page"`
Size int `json:"size"`
}
// TaskArchiveSnapshot is serialized into manifest.json inside the cold package.
// Small tables remain JSON arrays in Tables. Large v2 tables are streamed to
// package files listed in StreamedTables so their size is not bounded by memory.
type TaskArchiveSnapshot struct {
FormatVersion int `json:"format_version"`
CreatedAt time.Time `json:"created_at"`
TaskID int64 `json:"task_id"`
ExplorationID int64 `json:"exploration_id"`
SourceTaskIDs []int64 `json:"source_task_ids"`
Hosts []string `json:"hosts"`
ExclusiveHosts []string `json:"exclusive_hosts"`
ExclusiveAssetIDs []int64 `json:"exclusive_asset_ids"`
Tables map[string]json.RawMessage `json:"tables"`
StreamedTables map[string]string `json:"streamed_tables,omitempty"`
DataCounts map[string]int64 `json:"data_counts"`
AggregateStats map[string]any `json:"aggregate_stats"`
}
func scanTaskArchive(sc interface{ Scan(...any) error }) (*TaskArchive, error) {
var item TaskArchive
var sources string
err := sc.Scan(
&item.ID, &item.TaskID, &item.State, &item.Phase, &item.Progress, &item.Error,
&item.Warnings, &item.FormatVersion, &item.ArchivePath, &item.SHA256,
&item.OriginalSize, &item.CompressedSize, &item.TaskName, &item.TaskDescription,
&item.TaskGoal, &item.OriginalStatus, &item.CategoryIDSnapshot,
&item.CategoryNameSnapshot, &sources, &item.RemainingTimeoutSeconds,
&item.DataCounts, &item.AggregateStats, &item.ArchivedAt, &item.RequestedAt,
&item.CreatedAt, &item.UpdatedAt,
)
if err != nil {
return nil, err
}
if err := json.Unmarshal([]byte(sources), &item.SourceTaskIDs); err != nil {
return nil, err
}
if len(item.Warnings) == 0 {
item.Warnings = json.RawMessage("[]")
}
if len(item.DataCounts) == 0 {
item.DataCounts = json.RawMessage("{}")
}
if len(item.AggregateStats) == 0 {
item.AggregateStats = json.RawMessage("{}")
}
return &item, nil
}
const taskArchiveCols = `id, task_id, state, phase, progress, COALESCE(error,''), warnings,
format_version, COALESCE(archive_path,''), COALESCE(sha256,''), original_size,
compressed_size, COALESCE(task_name,''), COALESCE(task_description,''),
COALESCE(task_goal,''), COALESCE(original_status,''), category_id_snapshot,
COALESCE(category_name_snapshot,''), array_to_json(source_task_ids)::text,
remaining_timeout_seconds, data_counts, aggregate_stats, archived_at,
requested_at, created_at, updated_at`
func (d *DB) GetTaskArchive(id int64) (*TaskArchive, error) {
item, err := scanTaskArchive(d.QueryRow(`SELECT `+taskArchiveCols+` FROM task_archives WHERE id=$1`, id))
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
return item, err
}
func (d *DB) GetTaskArchiveByTask(taskID int64) (*TaskArchive, error) {
item, err := scanTaskArchive(d.QueryRow(`SELECT `+taskArchiveCols+` FROM task_archives WHERE task_id=$1`, taskID))
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
return item, err
}
func (d *DB) ListTaskArchives(search, state string, page, size int) (TaskArchivePage, error) {
if page < 1 {
page = 1
}
if size <= 0 || size > 100 {
size = 20
}
search = strings.TrimSpace(search)
state = strings.TrimSpace(state)
where := `WHERE ($1='' OR task_id::text ILIKE '%'||$1||'%' OR task_name ILIKE '%'||$1||'%' OR task_description ILIKE '%'||$1||'%')
AND ($2='' OR state=$2)`
var out TaskArchivePage
out.Page, out.Size = page, size
if err := d.QueryRow(`SELECT count(*) FROM task_archives `+where, search, state).Scan(&out.Total); err != nil {
return out, err
}
rows, err := d.Query(`SELECT `+taskArchiveCols+` FROM task_archives `+where+`
ORDER BY COALESCE(archived_at, requested_at) DESC, id DESC LIMIT $3 OFFSET $4`, search, state, size, (page-1)*size)
if err != nil {
return out, err
}
defer rows.Close()
for rows.Next() {
item, err := scanTaskArchive(rows)
if err != nil {
return out, err
}
out.Items = append(out.Items, *item)
}
return out, rows.Err()
}
// QueueTaskArchive validates lifecycle and direct inheritance while holding the
// task row. A failed archive can be explicitly retried through the same API.
func (d *DB) QueueTaskArchive(taskID int64) (*TaskArchive, error) {
tx, err := d.Begin()
if err != nil {
return nil, err
}
defer tx.Rollback() //nolint:errcheck
var name, description, goal, status, categoryName string
var categoryID *int64
var paused, queued bool
var deadline *time.Time
err = tx.QueryRow(`SELECT COALESCE(t.name,''), t.description, t.goal, t.status,
t.category_id, COALESCE(c.name,''), t.paused, t.queued, t.deadline_at
FROM tasks t LEFT JOIN task_categories c ON c.id=t.category_id
WHERE t.id=$1 AND t.deleted_at IS NULL FOR UPDATE OF t`, taskID).Scan(
&name, &description, &goal, &status, &categoryID, &categoryName, &paused, &queued, &deadline,
)
if errors.Is(err, sql.ErrNoRows) {
return nil, ErrTaskArchiveNotFound
}
if err != nil {
return nil, err
}
if queued {
return nil, ErrTaskArchiveQueued
}
if !paused && !IsTerminal(status) {
return nil, ErrTaskArchiveIneligible
}
var dependent int64
err = tx.QueryRow(`SELECT child.id FROM task_relations relation
JOIN tasks child ON child.id=relation.task_id AND child.deleted_at IS NULL
LEFT JOIN task_archives pending ON pending.task_id=child.id
WHERE relation.source_task_id=$1
AND (pending.id IS NULL OR pending.state NOT IN ('archive_queued','archiving'))
LIMIT 1`, taskID).Scan(&dependent)
if err == nil {
return nil, fmt.Errorf("%w: task %d", ErrTaskArchiveDependent, dependent)
}
if !errors.Is(err, sql.ErrNoRows) {
return nil, err
}
var sources []int64
rows, err := tx.Query(`SELECT source_task_id FROM task_relations WHERE task_id=$1 ORDER BY created_at, source_task_id`, taskID)
if err != nil {
return nil, err
}
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
rows.Close()
return nil, err
}
sources = append(sources, id)
}
if err := rows.Close(); err != nil {
return nil, err
}
if sources == nil {
sources = []int64{}
}
remaining := int64(0)
if paused && deadline != nil {
remaining = int64(time.Until(*deadline).Seconds())
if remaining < 0 {
remaining = 0
}
}
_, err = tx.Exec(`INSERT INTO task_archives(
task_id,state,phase,progress,error,warnings,format_version,task_name,
task_description,task_goal,original_status,category_id_snapshot,
category_name_snapshot,source_task_ids,remaining_timeout_seconds,requested_at)
VALUES ($1,$2,'queued',0,'','[]',$3,$4,$5,$6,$7,$8,$9,$10,$11,now())
ON CONFLICT (task_id) DO UPDATE SET
state=CASE WHEN task_archives.state IN ('archive_failed') THEN EXCLUDED.state ELSE task_archives.state END,
phase=CASE WHEN task_archives.state IN ('archive_failed') THEN 'queued' ELSE task_archives.phase END,
progress=CASE WHEN task_archives.state IN ('archive_failed') THEN 0 ELSE task_archives.progress END,
error=CASE WHEN task_archives.state IN ('archive_failed') THEN '' ELSE task_archives.error END,
format_version=CASE WHEN task_archives.state IN ('archive_failed') THEN EXCLUDED.format_version ELSE task_archives.format_version END,
requested_at=CASE WHEN task_archives.state IN ('archive_failed') THEN now() ELSE task_archives.requested_at END`,
taskID, ArchiveQueued, TaskArchiveFormatVersion, name, description, goal, status,
categoryID, categoryName, sources, remaining)
if err != nil {
return nil, err
}
item, err := scanTaskArchive(tx.QueryRow(`SELECT `+taskArchiveCols+` FROM task_archives WHERE task_id=$1`, taskID))
if err != nil {
return nil, err
}
if item.State != ArchiveQueued && item.State != ArchiveFailed {
return nil, fmt.Errorf("%w: current state %s", ErrTaskArchiveState, item.State)
}
return item, tx.Commit()
}
func (d *DB) QueueTaskArchiveRestore(id int64) (*TaskArchive, error) {
item, err := scanTaskArchive(d.QueryRow(`UPDATE task_archives
SET state=$2, phase='queued', progress=0, error='', requested_at=now()
WHERE id=$1 AND state IN ('ready','restore_failed') RETURNING `+taskArchiveCols, id, RestoreQueued))
if errors.Is(err, sql.ErrNoRows) {
return nil, ErrTaskArchiveState
}
return item, err
}
func (d *DB) QueueTaskArchiveDelete(id int64) (*TaskArchive, error) {
tx, err := d.Begin()
if err != nil {
return nil, err
}
defer tx.Rollback() //nolint:errcheck
var taskID int64
if err := tx.QueryRow(`SELECT task_id FROM task_archives WHERE id=$1 AND state IN ('ready','delete_failed') FOR UPDATE`, id).Scan(&taskID); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, ErrTaskArchiveState
}
return nil, err
}
var dependent int64
err = tx.QueryRow(`SELECT task_id FROM task_archives
WHERE id<>$1 AND $2=ANY(source_task_ids) AND state NOT IN ('delete_queued','deleting') LIMIT 1`, id, taskID).Scan(&dependent)
if err == nil {
return nil, fmt.Errorf("%w: task %d", ErrTaskArchiveDeleteBlocked, dependent)
}
if !errors.Is(err, sql.ErrNoRows) {
return nil, err
}
item, err := scanTaskArchive(tx.QueryRow(`UPDATE task_archives
SET state=$2, phase='queued', progress=0, error='', requested_at=now()
WHERE id=$1 RETURNING `+taskArchiveCols, id, DeleteQueued))
if err != nil {
return nil, err
}
return item, tx.Commit()
}
// RecoverTaskArchiveJobs keeps restore/delete resumable after an unclean shutdown.
// An interrupted archive requires an explicit retry: automatic startup retries can
// otherwise form a crash loop when the prior process was killed by resource limits.
func (d *DB) RecoverTaskArchiveJobs() error {
_, err := d.Exec(`UPDATE task_archives SET
state=CASE state WHEN 'archiving' THEN 'archive_failed'
WHEN 'restoring' THEN 'restore_queued'
WHEN 'deleting' THEN 'delete_queued' ELSE state END,
phase='interrupted',
error=CASE WHEN state='archiving' THEN '上次归档进程异常退出,请手动重试' ELSE '' END
WHERE state IN ('archiving','restoring','deleting')`)
return err
}
// ClaimTaskArchiveJob claims one persistent FIFO item for the single archive
// worker. It returns nil when the queue is empty.
func (d *DB) ClaimTaskArchiveJob(ctx context.Context) (*TaskArchive, error) {
tx, err := d.BeginTx(ctx, nil)
if err != nil {
return nil, err
}
defer tx.Rollback() //nolint:errcheck
var id int64
var queuedState string
err = tx.QueryRowContext(ctx, `SELECT id,state FROM task_archives
WHERE state IN ('archive_queued','restore_queued','delete_queued')
ORDER BY requested_at,id FOR UPDATE SKIP LOCKED LIMIT 1`).Scan(&id, &queuedState)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
if err != nil {
return nil, err
}
active := map[string]string{ArchiveQueued: Archiving, RestoreQueued: Restoring, DeleteQueued: Deleting}[queuedState]
item, err := scanTaskArchive(tx.QueryRowContext(ctx, `UPDATE task_archives
SET state=$2,phase='starting',progress=1,error='' WHERE id=$1 RETURNING `+taskArchiveCols, id, active))
if err != nil {
return nil, err
}
return item, tx.Commit()
}
func (d *DB) UpdateTaskArchiveProgress(id int64, phase string, progress int) error {
if progress < 0 {
progress = 0
}
if progress > 100 {
progress = 100
}
_, err := d.Exec(`UPDATE task_archives SET phase=$2,progress=$3 WHERE id=$1`, id, phase, progress)
return err
}
func (d *DB) AppendTaskArchiveWarning(id int64, warning string) error {
if strings.TrimSpace(warning) == "" {
return nil
}
_, err := d.Exec(`UPDATE task_archives SET warnings=warnings || jsonb_build_array($2::text) WHERE id=$1`, id, warning)
return err
}
func (d *DB) IsTaskArchiveRestored(id int64) (bool, error) {
var restored bool
err := d.QueryRow(`SELECT task.deleted_at IS NULL AND task.archived_at IS NULL
FROM task_archives archive JOIN tasks task ON task.id=archive.task_id WHERE archive.id=$1`, id).Scan(&restored)
if errors.Is(err, sql.ErrNoRows) {
return false, ErrTaskArchiveNotFound
}
return restored, err
}
func (d *DB) FailTaskArchiveJob(id int64, activeState string, cause error) error {
failed := map[string]string{Archiving: ArchiveFailed, Restoring: RestoreFailed, Deleting: DeleteFailed}[activeState]
if failed == "" {
return fmt.Errorf("unknown active archive state %q", activeState)
}
message := "unknown archive failure"
if cause != nil {
message = cause.Error()
}
_, err := d.Exec(`UPDATE task_archives SET state=$2,phase='failed',error=$3 WHERE id=$1`, id, failed, message)
return err
}
func queryArchiveRows(q interface {
Query(query string, args ...any) (*sql.Rows, error)
}, inner string, args ...any) (json.RawMessage, int64, error) {
// Do not aggregate the result in PostgreSQL. A jsonb array has a hard limit
// of 256 MiB for its elements, which large LLM request/response histories can
// exceed even though every individual record is valid. Reading row JSON in
// order also avoids building a second copy of the full table in PostgreSQL.
rows, err := q.Query(`SELECT row_to_json(row_data)::text FROM (`+inner+`) row_data`, args...)
if err != nil {
return nil, 0, err
}
defer rows.Close()
return encodeArchiveRows(rows)
}
func encodeArchiveRows(rows interface {
Next() bool
Scan(dest ...any) error
Err() error
}) (json.RawMessage, int64, error) {
var output bytes.Buffer
output.WriteByte('[')
var count int64
for rows.Next() {
var raw []byte
if err := rows.Scan(&raw); err != nil {
return nil, 0, err
}
if count > 0 {
output.WriteByte(',')
}
output.Write(raw)
count++
}
if err := rows.Err(); err != nil {
return nil, 0, err
}
output.WriteByte(']')
return json.RawMessage(output.Bytes()), count, nil
}
func writeArchiveRows(rows interface {
Next() bool
Scan(dest ...any) error
Err() error
}, writer io.Writer) (int64, error) {
var count int64
for rows.Next() {
var raw []byte
if err := rows.Scan(&raw); err != nil {
return 0, err
}
if written, err := writer.Write(raw); err != nil {
return 0, err
} else if written != len(raw) {
return 0, io.ErrShortWrite
}
if written, err := io.WriteString(writer, "\n"); err != nil {
return 0, err
} else if written != 1 {
return 0, io.ErrShortWrite
}
count++
}
if err := rows.Err(); err != nil {
return 0, err
}
return count, nil
}
func streamArchiveRows(q interface {
Query(query string, args ...any) (*sql.Rows, error)
}, writer io.Writer, inner string, args ...any) (int64, error) {
rows, err := q.Query(`SELECT row_to_json(row_data)::text FROM (`+inner+`) row_data`, args...)
if err != nil {
return 0, err
}
defer rows.Close()
return writeArchiveRows(rows, writer)
}
func rawRowCount(raw json.RawMessage) int64 {
var rows []json.RawMessage
if json.Unmarshal(raw, &rows) != nil {
return 0
}
return int64(len(rows))
}
func archiveAssetIDsQuery() string {
return `SELECT id FROM assets WHERE $1=ANY(task_ids)
UNION SELECT link.asset_id FROM task_asset_links link WHERE link.task_id=$1
UNION SELECT anchor.asset_id FROM exploration_anchors anchor
JOIN exploration_nodes node ON node.id=anchor.node_id WHERE node.exploration_id=$2
UNION SELECT value::bigint FROM findings finding
CROSS JOIN LATERAL jsonb_array_elements_text(
CASE WHEN jsonb_typeof(finding.asset_ids)='array' THEN finding.asset_ids ELSE '[]'::jsonb END
) value WHERE finding.task_id=$1 AND value ~ '^[0-9]+$'`
}
// SnapshotTaskArchive reads one repeatable PostgreSQL snapshot. Task-owned Agent
// writes are already quiescent at the server barrier; repeatable-read also keeps
// the asset and accounting views mutually consistent during serialization.
func (d *DB) SnapshotTaskArchive(taskID int64) (*TaskArchiveSnapshot, error) {
return d.snapshotTaskArchive(taskID, nil)
}
// SnapshotTaskArchiveWithLLMRecords streams the heavyweight record history to
// llmRecords while all other task-owned data is read from the same repeatable
// PostgreSQL snapshot.
func (d *DB) SnapshotTaskArchiveWithLLMRecords(taskID int64, llmRecords io.Writer) (*TaskArchiveSnapshot, error) {
if llmRecords == nil {
return nil, errors.New("nil LLM record archive writer")
}
return d.snapshotTaskArchive(taskID, llmRecords)
}
func (d *DB) snapshotTaskArchive(taskID int64, llmRecords io.Writer) (*TaskArchiveSnapshot, error) {
tx, err := d.BeginTx(context.Background(), &sql.TxOptions{Isolation: sql.LevelRepeatableRead, ReadOnly: true})
if err != nil {
return nil, err
}
defer tx.Rollback() //nolint:errcheck
if err := coordinateWithSchemaMigration(tx); err != nil {
return nil, err
}
var expID int64
if err := tx.QueryRow(`SELECT exploration_id FROM tasks WHERE id=$1 AND deleted_at IS NULL`, taskID).Scan(&expID); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, ErrTaskArchiveNotFound
}
return nil, err
}
tables := make(map[string]json.RawMessage)
queries := []struct {
name string
query string
args []any
}{
{"tasks", `SELECT * FROM tasks WHERE id=$1`, []any{taskID}},
{"explorations", `SELECT * FROM explorations WHERE id=$1`, []any{expID}},
{"exploration_nodes", `SELECT * FROM exploration_nodes WHERE exploration_id=$1 ORDER BY id`, []any{expID}},
{"exploration_edges", `SELECT * FROM exploration_edges WHERE exploration_id=$1 ORDER BY src_id,dst_id`, []any{expID}},
{"exploration_anchors", `SELECT anchor.* FROM exploration_anchors anchor JOIN exploration_nodes node ON node.id=anchor.node_id WHERE node.exploration_id=$1 ORDER BY node_id,asset_id`, []any{expID}},
{"task_constraints", `SELECT * FROM task_constraints WHERE exploration_id=$1 ORDER BY id`, []any{expID}},
{"activity", `SELECT * FROM activity WHERE exploration_id=$1 ORDER BY id`, []any{expID}},
{"task_relations", `SELECT * FROM task_relations WHERE task_id=$1 ORDER BY created_at,source_task_id`, []any{taskID}},
{"task_asset_links", `SELECT * FROM task_asset_links WHERE task_id=$1 ORDER BY asset_id`, []any{taskID}},
{"task_llm_profiles", `SELECT * FROM task_llm_profiles WHERE task_id=$1 ORDER BY position`, []any{taskID}},
{"task_scope", `SELECT * FROM task_scope WHERE task_id=$1 ORDER BY id`, []any{taskID}},
{"findings", `SELECT * FROM findings WHERE task_id=$1 ORDER BY id`, []any{taskID}},
{"finding_traffic_bindings", `SELECT b.* FROM finding_traffic_bindings b JOIN findings f ON f.id=b.finding_id WHERE f.task_id=$1 ORDER BY b.finding_id,b.position,b.id`, []any{taskID}},
{"traffic_evidence_snapshots", `SELECT s.* FROM traffic_evidence_snapshots s WHERE EXISTS(SELECT 1 FROM finding_traffic_bindings b JOIN findings f ON f.id=b.finding_id WHERE b.snapshot_id=s.id AND f.task_id=$1) ORDER BY s.id`, []any{taskID}},
{"llm_records", `SELECT * FROM llm_records WHERE COALESCE(task_id,'')=$1 ORDER BY id`, []any{strconv.FormatInt(taskID, 10)}},
{"llm_usage", `SELECT * FROM llm_usage WHERE COALESCE(task_id,'')=$1 OR exploration_id=$2 ORDER BY id`, []any{strconv.FormatInt(taskID, 10), expID}},
{"skill_usage", `SELECT * FROM skill_usage WHERE task_id=$1 OR exploration_id=$2 ORDER BY id`, []any{taskID, expID}},
{"tool_usage", `SELECT * FROM tool_usage WHERE task_id=$1 OR exploration_id=$2 ORDER BY id`, []any{taskID, expID}},
{"intercept_pending", `SELECT * FROM intercept_pending WHERE COALESCE(task_id,'')=$1 ORDER BY id`, []any{strconv.FormatInt(taskID, 10)}},
{"side_question_sessions", `SELECT * FROM side_question_sessions WHERE task_id=$1 ORDER BY session_key`, []any{taskID}},
{"side_question_requests", `SELECT r.* FROM side_question_requests r JOIN side_question_sessions s ON s.session_key=r.session_key WHERE s.task_id=$1 ORDER BY r.ordinal`, []any{taskID}},
{"assets", `SELECT asset.* FROM assets asset WHERE asset.id IN (` + archiveAssetIDsQuery() + `) ORDER BY asset.id`, []any{taskID, expID}},
}
counts := make(map[string]int64, len(queries))
streamedTables := map[string]string{}
for _, query := range queries {
if query.name == "llm_records" && llmRecords != nil {
count, err := streamArchiveRows(tx, llmRecords, query.query, query.args...)
if err != nil {
return nil, fmt.Errorf("snapshot %s: %w", query.name, err)
}
tables[query.name] = json.RawMessage("[]")
counts[query.name] = count
streamedTables[query.name] = TaskArchiveLLMRecordsPath
continue
}
raw, count, err := queryArchiveRows(tx, query.query, query.args...)
if err != nil {
return nil, fmt.Errorf("snapshot %s: %w", query.name, err)
}
tables[query.name] = raw
counts[query.name] = count
}
assetIDs, exclusiveAssetIDs, hosts, exclusiveHosts, err := archiveAssetMetadata(tx, taskID, expID, tables["assets"])
if err != nil {
return nil, err
}
_ = assetIDs // retained in the assets table payload; only exclusive ids need a side channel.
var sources []int64
rows, err := tx.Query(`SELECT source_task_id FROM task_relations WHERE task_id=$1 ORDER BY created_at,source_task_id`, taskID)
if err != nil {
return nil, err
}
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
rows.Close()
return nil, err
}
sources = append(sources, id)
}
if err := rows.Close(); err != nil {
return nil, err
}
stats, err := taskArchiveAggregates(tx, taskID, expID)
if err != nil {
return nil, err
}
snapshot := &TaskArchiveSnapshot{
FormatVersion: TaskArchiveFormatVersion, CreatedAt: time.Now().UTC(), TaskID: taskID,
ExplorationID: expID, SourceTaskIDs: sources, Hosts: hosts, ExclusiveHosts: exclusiveHosts,
ExclusiveAssetIDs: exclusiveAssetIDs, Tables: tables, StreamedTables: streamedTables,
DataCounts: counts, AggregateStats: stats,
}
return snapshot, tx.Commit()
}
func archiveAssetMetadata(tx *sql.Tx, taskID, expID int64, assetRows json.RawMessage) ([]int64, []int64, []string, []string, error) {
var rows []map[string]any
if err := json.Unmarshal(assetRows, &rows); err != nil {
return nil, nil, nil, nil, err
}
allHosts := map[string]struct{}{}
assetIDs := make([]int64, 0, len(rows))
for _, row := range rows {
if id, ok := jsonInt64(row["id"]); ok {
assetIDs = append(assetIDs, id)
}
for _, key := range []string{"domain", "ip"} {
if value, _ := row[key].(string); strings.TrimSpace(value) != "" {
allHosts[strings.ToLower(strings.TrimSpace(value))] = struct{}{}
}
}
if rawURL, _ := row["url"].(string); rawURL != "" {
if parsed, err := url.Parse(rawURL); err == nil && parsed.Hostname() != "" {
allHosts[strings.ToLower(parsed.Hostname())] = struct{}{}
}
}
}
exclusiveHosts, err := hostsForTaskDeletion(tx, taskID, expID)
if err != nil {
return nil, nil, nil, nil, err
}
rowsID, err := tx.Query(`WITH candidate AS (`+archiveAssetIDsQuery()+`)
SELECT asset.id FROM assets asset JOIN candidate ON candidate.id=asset.id
WHERE asset.company_id IS NULL
AND NOT EXISTS (SELECT 1 FROM tasks task WHERE task.id<>$1 AND task.deleted_at IS NULL AND task.id=ANY(asset.task_ids))
AND NOT EXISTS (
SELECT 1 FROM exploration_anchors anchor JOIN exploration_nodes node ON node.id=anchor.node_id
JOIN tasks task ON task.exploration_id=node.exploration_id
WHERE anchor.asset_id=asset.id AND task.id<>$1 AND task.deleted_at IS NULL
) ORDER BY asset.id`, taskID, expID)
if err != nil {
return nil, nil, nil, nil, err
}
var exclusiveIDs []int64
for rowsID.Next() {
var id int64
if err := rowsID.Scan(&id); err != nil {
rowsID.Close()
return nil, nil, nil, nil, err
}
exclusiveIDs = append(exclusiveIDs, id)
}
if err := rowsID.Close(); err != nil {
return nil, nil, nil, nil, err
}
hosts := make([]string, 0, len(allHosts))
for host := range allHosts {
hosts = append(hosts, host)
}
sort.Strings(hosts)
return assetIDs, exclusiveIDs, hosts, exclusiveHosts, nil
}
func taskArchiveAggregates(tx *sql.Tx, taskID, expID int64) (map[string]any, error) {
stats := map[string]any{}
var calls, input, output, cacheRead, cacheWrite int64
if err := tx.QueryRow(`SELECT count(*),COALESCE(sum(input_tokens),0),COALESCE(sum(output_tokens),0),
COALESCE(sum(cache_read),0),COALESCE(sum(cache_write),0)
FROM llm_usage WHERE COALESCE(task_id,'')=$1 OR exploration_id=$2`, strconv.FormatInt(taskID, 10), expID).
Scan(&calls, &input, &output, &cacheRead, &cacheWrite); err != nil {
return nil, err
}
stats["tokens"] = map[string]int64{"calls": calls, "input_tokens": input, "output_tokens": output, "cache_read_tokens": cacheRead, "cache_write_tokens": cacheWrite}
for _, item := range []struct {
name string
query string
args []any
}{
{"token_profiles", `SELECT COALESCE(jsonb_agg(to_jsonb(x)),'[]'::jsonb) FROM (
SELECT COALESCE(profile_name,'') profile_name,count(*) calls,1 tasks,
COALESCE(sum(input_tokens),0) input_tokens,COALESCE(sum(output_tokens),0) output_tokens,
COALESCE(sum(cache_read),0) cache_read_tokens,COALESCE(sum(cache_write),0) cache_write_tokens
FROM llm_usage WHERE COALESCE(task_id,'')=$1 OR exploration_id=$2
GROUP BY profile_name ORDER BY sum(input_tokens)+sum(output_tokens) DESC) x`, []any{strconv.FormatInt(taskID, 10), expID}},
{"token_daily", `SELECT COALESCE(jsonb_agg(to_jsonb(x)),'[]'::jsonb) FROM (
SELECT COALESCE(profile_name,'') profile_name,to_char(ts AT TIME ZONE 'UTC','YYYY-MM-DD') date,
COALESCE(sum(input_tokens),0) input_tokens,COALESCE(sum(output_tokens),0) output_tokens,
COALESCE(sum(cache_read),0) cache_read_tokens
FROM llm_usage WHERE COALESCE(task_id,'')=$1 OR exploration_id=$2
GROUP BY profile_name,date ORDER BY date) x`, []any{strconv.FormatInt(taskID, 10), expID}},
{"skills", `SELECT COALESCE(jsonb_object_agg(name,n),'{}'::jsonb) FROM (SELECT skill name,count(*) n FROM skill_usage WHERE (task_id=$1 OR exploration_id=$2) AND found GROUP BY skill) x`, []any{taskID, expID}},
{"skill_stats", `SELECT COALESCE(jsonb_agg(to_jsonb(x)),'[]'::jsonb) FROM (
SELECT skill,count(*) calls,1 tasks,
COALESCE(array_agg(DISTINCT agent_key) FILTER (WHERE agent_key IS NOT NULL),ARRAY[]::text[]) agents,
max(ts) last_used
FROM skill_usage WHERE (task_id=$1 OR exploration_id=$2) AND found GROUP BY skill) x`, []any{taskID, expID}},
{"missing_skill_stats", `SELECT COALESCE(jsonb_agg(to_jsonb(x)),'[]'::jsonb) FROM (
SELECT skill,count(*) calls,0 tasks,
COALESCE(array_agg(DISTINCT agent_key) FILTER (WHERE agent_key IS NOT NULL),ARRAY[]::text[]) agents,
max(ts) last_used
FROM skill_usage WHERE (task_id=$1 OR exploration_id=$2) AND NOT found GROUP BY skill) x`, []any{taskID, expID}},
{"tools", `SELECT COALESCE(jsonb_object_agg(name,n),'{}'::jsonb) FROM (SELECT tool_key name,count(*) n FROM tool_usage WHERE task_id=$1 OR exploration_id=$2 GROUP BY tool_key) x`, []any{taskID, expID}},
{"findings", `SELECT COALESCE(jsonb_object_agg(name,n),'{}'::jsonb) FROM (SELECT COALESCE(NULLIF(severity,''),'unknown') name,count(*) n FROM findings WHERE task_id=$1 GROUP BY severity) x`, []any{taskID}},
{"finding_stats", `SELECT jsonb_build_object(
'total',count(*),'pending',count(*) FILTER (WHERE status='pending'),
'critical',count(*) FILTER (WHERE severity='critical'),'high',count(*) FILTER (WHERE severity='high'),
'medium',count(*) FILTER (WHERE severity='medium'),'low',count(*) FILTER (WHERE severity='low'),
'vulnclasses',COALESCE(jsonb_agg(DISTINCT vulnclass) FILTER (WHERE vulnclass<>''),'[]'::jsonb))
FROM findings WHERE task_id=$1`, []any{taskID}},
} {
var raw []byte
if err := tx.QueryRow(item.query, item.args...).Scan(&raw); err != nil {
return nil, err
}
var value any
if err := json.Unmarshal(raw, &value); err != nil {
return nil, err
}
stats[item.name] = value
}
return stats, nil
}
func jsonInt64(value any) (int64, bool) {
switch value := value.(type) {
case float64:
return int64(value), value == float64(int64(value))
case json.Number:
id, err := value.Int64()
return id, err == nil
case string:
id, err := strconv.ParseInt(value, 10, 64)
return id, err == nil
default:
return 0, false
}
}
+793
View File
@@ -0,0 +1,793 @@
package db
import (
"bytes"
"database/sql"
"encoding/json"
"errors"
"fmt"
"io"
"strconv"
"strings"
"time"
)
// CompleteTaskArchive performs the hot-store compaction only after the external
// package has been fully written and checksummed. The task/exploration rows remain
// as minimal ID stubs; all heavyweight task-owned rows move into the package.
func (d *DB) CompleteTaskArchive(
archiveID int64,
snapshot *TaskArchiveSnapshot,
archivePath, sha256 string,
originalSize, compressedSize int64,
) error {
if snapshot == nil || snapshot.FormatVersion != TaskArchiveFormatVersion {
return ErrTaskArchiveFormatMismatch
}
countsRaw, err := json.Marshal(snapshot.DataCounts)
if err != nil {
return err
}
statsRaw, err := json.Marshal(snapshot.AggregateStats)
if err != nil {
return err
}
tx, err := d.Begin()
if err != nil {
return err
}
defer tx.Rollback() //nolint:errcheck
if err := coordinateWithSchemaMigration(tx); err != nil {
return err
}
var taskID, expID int64
var state string
if err := tx.QueryRow(`SELECT archive.task_id,task.exploration_id,archive.state
FROM task_archives archive JOIN tasks task ON task.id=archive.task_id
WHERE archive.id=$1 FOR UPDATE OF archive,task`, archiveID).Scan(&taskID, &expID, &state); err != nil {
return err
}
if state != Archiving || taskID != snapshot.TaskID || expID != snapshot.ExplorationID {
return fmt.Errorf("%w: archive snapshot identity/state mismatch", ErrTaskArchiveState)
}
var dependent int64
err = tx.QueryRow(`SELECT child.id FROM task_relations relation
JOIN tasks child ON child.id=relation.task_id AND child.deleted_at IS NULL
WHERE relation.source_task_id=$1 LIMIT 1`, taskID).Scan(&dependent)
if err == nil {
return fmt.Errorf("%w: task %d", ErrTaskArchiveDependent, dependent)
}
if !errors.Is(err, sql.ErrNoRows) {
return err
}
// Prevent task/asset ownership from changing while exclusivity is rechecked.
if _, err := tx.Exec(`LOCK TABLE assets, exploration_anchors IN SHARE ROW EXCLUSIVE MODE`); err != nil {
return err
}
if _, err := tx.Exec(`DELETE FROM intercept_pending WHERE COALESCE(task_id,'')=$1`, strconv.FormatInt(taskID, 10)); err != nil {
return err
}
if _, err := tx.Exec(`DELETE FROM llm_records WHERE COALESCE(task_id,'')=$1`, strconv.FormatInt(taskID, 10)); err != nil {
return err
}
if _, err := tx.Exec(`DELETE FROM llm_usage WHERE COALESCE(task_id,'')=$1 OR exploration_id=$2`, strconv.FormatInt(taskID, 10), expID); err != nil {
return err
}
for _, table := range []string{"skill_usage", "tool_usage"} {
if _, err := tx.Exec(`DELETE FROM `+table+` WHERE task_id=$1 OR exploration_id=$2`, taskID, expID); err != nil {
return err
}
}
if _, err := tx.Exec(`DELETE FROM side_question_sessions WHERE task_id=$1`, taskID); err != nil {
return err
}
if _, err := tx.Exec(`DELETE FROM findings WHERE task_id=$1`, taskID); err != nil {
return err
}
for _, statement := range []struct {
query string
args []any
}{
{`DELETE FROM task_relations WHERE task_id=$1 OR source_task_id=$1`, []any{taskID}},
{`DELETE FROM task_asset_links WHERE task_id=$1`, []any{taskID}},
{`DELETE FROM task_llm_profiles WHERE task_id=$1`, []any{taskID}},
{`DELETE FROM task_scope WHERE task_id=$1`, []any{taskID}},
{`DELETE FROM task_constraints WHERE exploration_id=$1`, []any{expID}},
{`DELETE FROM activity WHERE exploration_id=$1`, []any{expID}},
{`DELETE FROM exploration_nodes WHERE exploration_id=$1`, []any{expID}},
} {
if _, err := tx.Exec(statement.query, statement.args...); err != nil {
return err
}
}
if _, err := tx.Exec(`UPDATE assets SET task_ids=array_remove(task_ids,$1) WHERE $1=ANY(task_ids)`, taskID); err != nil {
return err
}
if len(snapshot.ExclusiveAssetIDs) > 0 {
if _, err := tx.Exec(`DELETE FROM assets asset
WHERE asset.id=ANY($2::bigint[]) AND asset.company_id IS NULL
AND NOT EXISTS (SELECT 1 FROM tasks task WHERE task.id<>$1 AND task.deleted_at IS NULL AND task.id=ANY(asset.task_ids))
AND NOT EXISTS (
SELECT 1 FROM exploration_anchors anchor JOIN exploration_nodes node ON node.id=anchor.node_id
JOIN tasks task ON task.exploration_id=node.exploration_id
WHERE anchor.asset_id=asset.id AND task.id<>$1 AND task.deleted_at IS NULL
)`, taskID, snapshot.ExclusiveAssetIDs); err != nil {
return err
}
}
if _, err := tx.Exec(`UPDATE explorations SET description='',goal='',status='open' WHERE id=$1`, expID); err != nil {
return err
}
if _, err := tx.Exec(`UPDATE tasks SET
name='',category_id=NULL,description='',goal='',paused=true,queued=false,queued_at=NULL,queue_mode='',
llm_profile_id=NULL,active_llm_profile_id=NULL,llm_chain_revision=llm_chain_revision+1,
company_id=NULL,parent_ref=NULL,timeout_seconds=0,coverage_enabled=true,pinned_at=NULL,
first_run_at=NULL,deadline_at=NULL,archived_at=now(),deleted_at=now()
WHERE id=$1`, taskID); err != nil {
return err
}
if _, err := tx.Exec(`UPDATE task_archives SET
state=$2,phase='ready',progress=100,error='',archive_path=$3,sha256=$4,
original_size=$5,compressed_size=$6,data_counts=$7,aggregate_stats=$8,
format_version=$9,archived_at=now(),warnings='[]'
WHERE id=$1`, archiveID, ArchiveReady, archivePath, sha256, originalSize, compressedSize,
string(countsRaw), string(statsRaw), snapshot.FormatVersion); err != nil {
return err
}
return tx.Commit()
}
// RestoreTaskArchive restores PostgreSQL rows from a verified manifest. It is
// idempotent for accounting/traffic retry scenarios and returns non-fatal
// warnings for global objects that intentionally are not recreated.
func (d *DB) RestoreTaskArchive(archiveID int64, snapshot *TaskArchiveSnapshot, remainingTimeoutSeconds int64) ([]string, error) {
return d.restoreTaskArchive(archiveID, snapshot, remainingTimeoutSeconds, nil)
}
// RestoreTaskArchiveWithLLMRecords restores a v2 package whose heavyweight LLM
// record history is stored as a sequence of JSON objects outside manifest.json.
func (d *DB) RestoreTaskArchiveWithLLMRecords(
archiveID int64,
snapshot *TaskArchiveSnapshot,
remainingTimeoutSeconds int64,
llmRecords io.Reader,
) ([]string, error) {
if llmRecords == nil {
return nil, errors.New("nil streamed LLM record reader")
}
return d.restoreTaskArchive(archiveID, snapshot, remainingTimeoutSeconds, llmRecords)
}
func (d *DB) restoreTaskArchive(
archiveID int64,
snapshot *TaskArchiveSnapshot,
remainingTimeoutSeconds int64,
llmRecords io.Reader,
) ([]string, error) {
if snapshot == nil || !IsTaskArchiveFormatSupported(snapshot.FormatVersion) {
return nil, ErrTaskArchiveFormatMismatch
}
streamedLLMRecords := snapshot.StreamedTables["llm_records"]
if (len(snapshot.StreamedTables) > 0 && snapshot.FormatVersion < 2) ||
len(snapshot.StreamedTables) > 1 ||
(len(snapshot.StreamedTables) == 1 && streamedLLMRecords != TaskArchiveLLMRecordsPath) {
return nil, fmt.Errorf("%w: unsupported streamed table metadata", ErrTaskArchiveFormatMismatch)
}
if streamedLLMRecords != "" && llmRecords == nil {
return nil, fmt.Errorf("%w: streamed LLM records are missing", ErrTaskArchiveFormatMismatch)
}
tx, err := d.Begin()
if err != nil {
return nil, err
}
defer tx.Rollback() //nolint:errcheck
if err := coordinateWithSchemaMigration(tx); err != nil {
return nil, err
}
var taskID, expID int64
var state string
if err := tx.QueryRow(`SELECT archive.task_id,task.exploration_id,archive.state
FROM task_archives archive JOIN tasks task ON task.id=archive.task_id
WHERE archive.id=$1 FOR UPDATE OF archive,task`, archiveID).Scan(&taskID, &expID, &state); err != nil {
return nil, err
}
if state != Restoring || taskID != snapshot.TaskID || expID != snapshot.ExplorationID {
return nil, fmt.Errorf("%w: restore snapshot identity/state mismatch", ErrTaskArchiveState)
}
warnings := []string{}
assetMap, assetWarnings, err := restoreArchiveAssets(tx, taskID, snapshot.Tables["assets"])
if err != nil {
return nil, err
}
warnings = append(warnings, assetWarnings...)
// The task row exists as an archived stub. Restore its global references only
// when the current instance still owns them; never recreate categories/profiles.
taskRows, err := decodeArchiveRows(snapshot.Tables["tasks"])
if err != nil || len(taskRows) != 1 {
return nil, fmt.Errorf("restore task row: expected one row: %w", err)
}
taskRow := taskRows[0]
categoryID, _ := jsonInt64(taskRow["category_id"])
if categoryID > 0 && !rowExists(tx, "task_categories", categoryID) {
taskRow["category_id"] = nil
warnings = append(warnings, fmt.Sprintf("任务分类 %d 已删除,已恢复为未分类", categoryID))
}
companyID, _ := jsonInt64(taskRow["company_id"])
if companyID > 0 && !rowExists(tx, "companies", companyID) {
taskRow["company_id"] = nil
warnings = append(warnings, fmt.Sprintf("任务企业 %d 已删除,企业关联已跳过", companyID))
}
for _, key := range []string{"llm_profile_id", "active_llm_profile_id"} {
profileID, _ := jsonInt64(taskRow[key])
if profileID > 0 && !rowExists(tx, "llm_profiles", profileID) {
taskRow[key] = nil
warnings = append(warnings, fmt.Sprintf("LLM 配置 %d 已删除,已从任务配置中移除", profileID))
}
}
if err := restoreExplorationStub(tx, snapshot.Tables["explorations"], expID); err != nil {
return nil, err
}
if err := restoreTaskStub(tx, taskRow, taskID, remainingTimeoutSeconds); err != nil {
return nil, err
}
remappedTables, err := remapArchiveAssetReferences(snapshot.Tables, assetMap)
if err != nil {
return nil, err
}
remappedTables["findings"], err = normalizeArchivedFindingVersions(remappedTables["findings"])
if err != nil {
return nil, err
}
// Insert graph rows in foreign-key order. The archived stub has no graph rows,
// so an ID conflict signals external corruption and must stop the restore.
for _, table := range []string{"exploration_nodes", "exploration_edges", "exploration_anchors", "task_constraints", "activity"} {
if err := insertArchiveRows(tx, table, remappedTables[table]); err != nil {
return nil, fmt.Errorf("restore %s: %w", table, err)
}
}
if warning, err := restoreTaskRelations(tx, taskID, remappedTables["task_relations"]); err != nil {
return nil, err
} else {
warnings = append(warnings, warning...)
}
for _, table := range []string{"side_question_sessions", "side_question_requests"} {
if err := insertArchiveRows(tx, table, remappedTables[table]); err != nil {
return nil, fmt.Errorf("restore %s: %w", table, err)
}
}
if warning, err := restoreTaskScopes(tx, remappedTables["task_scope"]); err != nil {
return nil, err
} else {
warnings = append(warnings, warning...)
}
if warning, err := restoreTaskLLMProfiles(tx, remappedTables["task_llm_profiles"]); err != nil {
return nil, err
} else {
warnings = append(warnings, warning...)
}
for _, table := range []string{"task_asset_links", "findings"} {
if err := insertArchiveRows(tx, table, remappedTables[table]); err != nil {
return nil, fmt.Errorf("restore %s: %w", table, err)
}
}
if err := restoreFindingTrafficTx(tx, snapshot); err != nil {
return nil, fmt.Errorf("restore finding traffic: %w", err)
}
if streamedLLMRecords != "" {
count, err := insertArchiveJSONSequenceRows(tx, "llm_records", llmRecords)
if err != nil {
return nil, fmt.Errorf("restore llm_records: %w", err)
}
if expected := snapshot.DataCounts["llm_records"]; count != expected {
return nil, fmt.Errorf("restore llm_records: row count %d does not match manifest %d", count, expected)
}
} else if err := insertArchiveRows(tx, "llm_records", remappedTables["llm_records"]); err != nil {
return nil, fmt.Errorf("restore llm_records: %w", err)
}
for _, table := range []string{"llm_usage", "skill_usage", "tool_usage"} {
if err := insertArchiveRows(tx, table, remappedTables[table]); err != nil {
return nil, fmt.Errorf("restore %s: %w", table, err)
}
}
if err := restoreInterceptRows(tx, remappedTables["intercept_pending"]); err != nil {
return nil, err
}
warningsRaw, _ := json.Marshal(warnings)
if _, err := tx.Exec(`UPDATE task_archives SET warnings=$2,phase='database_restored',progress=85,error='' WHERE id=$1`, archiveID, string(warningsRaw)); err != nil {
return nil, err
}
return warnings, tx.Commit()
}
func restoreExplorationStub(tx *sql.Tx, raw json.RawMessage, expID int64) error {
_, err := tx.Exec(`UPDATE explorations current SET
description=archived.description,goal=archived.goal,status=archived.status,
created_at=archived.created_at,updated_at=archived.updated_at
FROM json_populate_record(NULL::explorations,$2::json) archived
WHERE current.id=$1 AND archived.id=$1`, expID, string(firstArchiveRow(raw)))
return err
}
func restoreTaskStub(tx *sql.Tx, row map[string]any, taskID, remaining int64) error {
raw, err := json.Marshal(row)
if err != nil {
return err
}
_, err = tx.Exec(`UPDATE tasks current SET
name=archived.name,category_id=archived.category_id,description=archived.description,goal=archived.goal,
status=archived.status,paused=archived.paused,queued=false,queued_at=NULL,queue_mode='',
llm_profile_id=archived.llm_profile_id,active_llm_profile_id=archived.active_llm_profile_id,
llm_chain_revision=archived.llm_chain_revision,company_id=archived.company_id,parent_ref=archived.parent_ref,
timeout_seconds=archived.timeout_seconds,plan_heartbeat_seconds=archived.plan_heartbeat_seconds,
coverage_enabled=archived.coverage_enabled,pinned_at=archived.pinned_at,first_run_at=archived.first_run_at,
deadline_at=CASE WHEN archived.paused AND $3>0 THEN now()+make_interval(secs=>$3::double precision)
ELSE archived.deadline_at END,
deleted_at=NULL,archived_at=NULL,completed_at=archived.completed_at,
created_at=archived.created_at,updated_at=archived.updated_at
FROM json_populate_record(NULL::tasks,$2::json) archived
WHERE current.id=$1 AND archived.id=$1`, taskID, string(raw), remaining)
return err
}
func firstArchiveRow(raw json.RawMessage) json.RawMessage {
var rows []json.RawMessage
if json.Unmarshal(raw, &rows) != nil || len(rows) == 0 {
return json.RawMessage("{}")
}
return rows[0]
}
func decodeArchiveRows(raw json.RawMessage) ([]map[string]any, error) {
if len(raw) == 0 {
return nil, nil
}
decoder := json.NewDecoder(bytes.NewReader(raw))
decoder.UseNumber()
var rows []map[string]any
if err := decoder.Decode(&rows); err != nil {
return nil, err
}
return rows, nil
}
func rowExists(tx *sql.Tx, table string, id int64) bool {
if table != "task_categories" && table != "llm_profiles" && table != "companies" && table != "tasks" && table != "assets" {
return false
}
var exists bool
_ = tx.QueryRow(`SELECT EXISTS(SELECT 1 FROM `+table+` WHERE id=$1)`, id).Scan(&exists)
return exists
}
func insertArchiveRows(tx *sql.Tx, table string, raw json.RawMessage) error {
// json_populate_recordset inserts NULL for absent columns, bypassing SQL
// defaults. Preserve compatibility with v3 archives predating side memory.
if (table == "side_question_sessions" || table == "side_question_requests") && len(raw) > 0 {
var rows []map[string]json.RawMessage
if err := json.Unmarshal(raw, &rows); err != nil {
return err
}
field := "memory"
if table == "side_question_requests" {
field = "context_info"
}
for _, row := range rows {
if len(row[field]) == 0 || string(row[field]) == "null" {
row[field] = json.RawMessage(`{}`)
}
}
var err error
raw, err = json.Marshal(rows)
if err != nil {
return err
}
}
allowed := map[string]bool{
"exploration_nodes": true, "exploration_edges": true, "exploration_anchors": true,
"task_constraints": true, "activity": true, "task_asset_links": true, "findings": true,
"llm_records": true, "llm_usage": true, "skill_usage": true, "tool_usage": true,
"side_question_sessions": true, "side_question_requests": true,
}
if !allowed[table] {
return fmt.Errorf("archive restore table %q is not allowed", table)
}
if rawRowCount(raw) == 0 {
return nil
}
_, err := tx.Exec(`INSERT INTO `+table+` SELECT * FROM json_populate_recordset(NULL::`+table+`,$1::json)`, string(raw))
return err
}
func insertArchiveJSONSequenceRows(tx *sql.Tx, table string, reader io.Reader) (int64, error) {
if table != "llm_records" {
return 0, fmt.Errorf("archive streamed restore table %q is not allowed", table)
}
statement, err := tx.Prepare(`INSERT INTO llm_records SELECT * FROM json_populate_record(NULL::llm_records,$1::json)`)
if err != nil {
return 0, err
}
defer statement.Close()
decoder := json.NewDecoder(reader)
var count int64
for {
var raw json.RawMessage
if err := decoder.Decode(&raw); errors.Is(err, io.EOF) {
break
} else if err != nil {
return 0, err
}
trimmed := bytes.TrimSpace(raw)
if len(trimmed) == 0 || trimmed[0] != '{' {
return 0, errors.New("streamed archive row must be a JSON object")
}
if _, err := statement.Exec(string(trimmed)); err != nil {
return 0, err
}
count++
}
return count, nil
}
func restoreArchiveAssets(tx *sql.Tx, taskID int64, raw json.RawMessage) (map[int64]int64, []string, error) {
rows, err := decodeArchiveRows(raw)
if err != nil {
return nil, nil, err
}
mapping := make(map[int64]int64, len(rows))
warnings := []string{}
for _, row := range rows {
oldID, ok := jsonInt64(row["id"])
if !ok || oldID <= 0 {
return nil, nil, errors.New("archived asset has invalid id")
}
if companyID, ok := jsonInt64(row["company_id"]); ok && companyID > 0 && !rowExists(tx, "companies", companyID) {
row["company_id"] = nil
row["company_source"] = "explicit"
warnings = append(warnings, fmt.Sprintf("资产 %d 的企业 %d 已删除,已恢复为未归属", oldID, companyID))
}
if existing, found, err := findArchiveAssetNaturalID(tx, row); err != nil {
return nil, nil, err
} else if found {
mapping[oldID] = existing
if _, err := tx.Exec(`UPDATE assets SET task_ids=CASE WHEN $1=ANY(task_ids) THEN task_ids ELSE array_append(task_ids,$1) END WHERE id=$2`, taskID, existing); err != nil {
return nil, nil, err
}
continue
}
candidate := oldID
if rowExists(tx, "assets", oldID) {
if err := tx.QueryRow(`SELECT nextval(pg_get_serial_sequence('assets','id'))`).Scan(&candidate); err != nil {
return nil, nil, err
}
row["id"] = candidate
warnings = append(warnings, fmt.Sprintf("资产 ID %d 已被占用,恢复为 %d", oldID, candidate))
}
row["task_ids"] = mergeJSONTaskID(row["task_ids"], taskID)
assetRaw, _ := json.Marshal(row)
res, err := tx.Exec(`INSERT INTO assets SELECT * FROM json_populate_record(NULL::assets,$1::json) ON CONFLICT DO NOTHING`, string(assetRaw))
if err != nil {
return nil, nil, err
}
inserted, _ := res.RowsAffected()
if inserted == 0 {
existing, found, findErr := findArchiveAssetNaturalID(tx, row)
if findErr != nil || !found {
return nil, nil, errors.Join(findErr, fmt.Errorf("could not restore asset %d", oldID))
}
candidate = existing
if _, err := tx.Exec(`UPDATE assets SET task_ids=CASE WHEN $1=ANY(task_ids) THEN task_ids ELSE array_append(task_ids,$1) END WHERE id=$2`, taskID, existing); err != nil {
return nil, nil, err
}
}
mapping[oldID] = candidate
}
return mapping, warnings, nil
}
func mergeJSONTaskID(value any, taskID int64) []int64 {
out := []int64{}
seen := map[int64]bool{}
if values, ok := value.([]any); ok {
for _, item := range values {
if id, ok := jsonInt64(item); ok && id > 0 && !seen[id] {
seen[id] = true
out = append(out, id)
}
}
}
if !seen[taskID] {
out = append(out, taskID)
}
return out
}
func findArchiveAssetNaturalID(tx *sql.Tx, row map[string]any) (int64, bool, error) {
typeName, _ := row["type"].(string)
stringValue := func(key string) string {
value, _ := row[key].(string)
return value
}
var query string
var args []any
switch typeName {
case "root_domain":
query, args = `SELECT id FROM assets WHERE type='root_domain' AND domain=$1`, []any{stringValue("domain")}
case "ip":
query, args = `SELECT id FROM assets WHERE type='ip' AND ip=$1`, []any{stringValue("ip")}
case "subdomain":
query, args = `SELECT id FROM assets WHERE type='subdomain' AND domain=$1 AND COALESCE(record_type,'')=$2`, []any{stringValue("domain"), stringValue("record_type")}
case "app":
if bundle := stringValue("bundle_id"); bundle != "" {
query, args = `SELECT id FROM assets WHERE type='app' AND bundle_id=$1`, []any{bundle}
} else {
query, args = `SELECT id FROM assets WHERE type='app' AND bundle_id IS NULL AND app_name=$1`, []any{stringValue("app_name")}
}
case "service":
if stringValue("service_type") == "http" {
query, args = `SELECT id FROM assets WHERE type='service' AND service_type='http' AND url=$1`, []any{stringValue("url")}
} else {
port, _ := jsonInt64(row["port"])
query, args = `SELECT id FROM assets WHERE type='service' AND service_type='other' AND COALESCE(domain,'')=$1 AND COALESCE(ip,'')=$2 AND port=$3 AND service_name=$4`, []any{stringValue("domain"), stringValue("ip"), port, stringValue("service_name")}
}
case "endpoint":
query, args = `SELECT id FROM assets WHERE type='endpoint' AND url=$1 AND method=$2`, []any{stringValue("url"), stringValue("method")}
default:
return 0, false, fmt.Errorf("unsupported archived asset type %q", typeName)
}
var id int64
err := tx.QueryRow(query, args...).Scan(&id)
if errors.Is(err, sql.ErrNoRows) {
return 0, false, nil
}
return id, err == nil, err
}
func remapArchiveAssetReferences(tables map[string]json.RawMessage, mapping map[int64]int64) (map[string]json.RawMessage, error) {
out := make(map[string]json.RawMessage, len(tables))
for name, raw := range tables {
out[name] = raw
}
for _, table := range []string{"exploration_anchors", "task_asset_links"} {
rows, err := decodeArchiveRows(tables[table])
if err != nil {
return nil, err
}
for _, row := range rows {
if old, ok := jsonInt64(row["asset_id"]); ok {
if replacement, exists := mapping[old]; exists {
row["asset_id"] = replacement
}
}
}
out[table], _ = json.Marshal(rows)
}
for _, table := range []string{"exploration_nodes", "findings"} {
rows, err := decodeArchiveRows(tables[table])
if err != nil {
return nil, err
}
for _, row := range rows {
key := "asset_ids"
container := row
if table == "exploration_nodes" {
payload, ok := row["payload"].(map[string]any)
if !ok {
continue
}
container = payload
}
if values, ok := container[key].([]any); ok {
for i, value := range values {
if old, ok := jsonInt64(value); ok {
if replacement, exists := mapping[old]; exists {
values[i] = replacement
}
}
}
container[key] = values
}
}
out[table], _ = json.Marshal(rows)
}
return out, nil
}
func restoreTaskRelations(tx *sql.Tx, taskID int64, raw json.RawMessage) ([]string, error) {
rows, err := decodeArchiveRows(raw)
if err != nil {
return nil, err
}
warnings := []string{}
for _, row := range rows {
sourceID, ok := jsonInt64(row["source_task_id"])
if !ok || !liveTaskExists(tx, sourceID) {
warnings = append(warnings, fmt.Sprintf("来源任务 %d 不可用,继承关系已跳过", sourceID))
continue
}
created, _ := row["created_at"].(string)
if _, err := tx.Exec(`INSERT INTO task_relations(task_id,source_task_id,created_at) VALUES($1,$2,COALESCE($3::timestamptz,now())) ON CONFLICT DO NOTHING`, taskID, sourceID, nilIfEmptyString(created)); err != nil {
return nil, err
}
}
return warnings, nil
}
func restoreTaskScopes(tx *sql.Tx, raw json.RawMessage) ([]string, error) {
rows, err := decodeArchiveRows(raw)
if err != nil {
return nil, err
}
warnings := []string{}
kept := make([]map[string]any, 0, len(rows))
for _, row := range rows {
if companyID, ok := jsonInt64(row["company_id"]); ok && companyID > 0 && !rowExists(tx, "companies", companyID) {
warnings = append(warnings, fmt.Sprintf("企业 %d 已删除,关联范围已跳过", companyID))
continue
}
kept = append(kept, row)
}
if len(kept) == 0 {
return warnings, nil
}
encoded, _ := json.Marshal(kept)
_, err = tx.Exec(`INSERT INTO task_scope SELECT * FROM json_populate_recordset(NULL::task_scope,$1::json) ON CONFLICT DO NOTHING`, string(encoded))
return warnings, err
}
func restoreTaskLLMProfiles(tx *sql.Tx, raw json.RawMessage) ([]string, error) {
rows, err := decodeArchiveRows(raw)
if err != nil {
return nil, err
}
warnings := []string{}
position := 0
for _, row := range rows {
profileID, ok := jsonInt64(row["profile_id"])
if !ok || !rowExists(tx, "llm_profiles", profileID) {
warnings = append(warnings, fmt.Sprintf("LLM 配置 %d 已删除,配置链项已跳过", profileID))
continue
}
taskID, _ := jsonInt64(row["task_id"])
status, _ := row["status"].(string)
lastError, _ := row["last_error"].(string)
exhaustedAt, _ := row["exhausted_at"].(string)
createdAt, _ := row["created_at"].(string)
updatedAt, _ := row["updated_at"].(string)
if _, err := tx.Exec(`INSERT INTO task_llm_profiles(task_id,profile_id,position,status,last_error,exhausted_at,created_at,updated_at)
VALUES($1,$2,$3,$4,NULLIF($5,''),$6::timestamptz,COALESCE($7::timestamptz,now()),COALESCE($8::timestamptz,now()))
ON CONFLICT DO NOTHING`, taskID, profileID, position, status, lastError,
nilIfEmptyString(exhaustedAt), nilIfEmptyString(createdAt), nilIfEmptyString(updatedAt)); err != nil {
return nil, err
}
position++
}
return warnings, nil
}
func restoreInterceptRows(tx *sql.Tx, raw json.RawMessage) error {
rows, err := decodeArchiveRows(raw)
if err != nil {
return err
}
for _, row := range rows {
// Archives predating approval snapshots lack this NOT NULL column.
if source, _ := row["decision_source"].(string); source == "" {
source = "unknown"
if row["rule_id"] != nil {
source = "rule"
} else if reason, _ := row["reason"].(string); strings.HasPrefix(reason, "[模型]") {
source = "model"
}
row["decision_source"] = source
}
if status, _ := row["status"].(string); status == "pending" {
row["status"] = "timeout"
row["reason"] = "任务归档期间审批已超时"
row["decided_at"] = time.Now().UTC()
if audit, ok := row["audit"].(map[string]any); ok {
audit["effective_action"] = "deny"
audit["decision_reason"] = "任务归档期间审批已超时"
audit["execution_status"] = "not_executed"
}
} else if audit, ok := row["audit"].(map[string]any); ok && audit["execution_status"] == "awaiting_result" {
// An archived run cannot resume its former result callback.
audit["execution_status"] = "unknown"
}
}
if len(rows) == 0 {
return nil
}
encoded, _ := json.Marshal(rows)
_, err = tx.Exec(`INSERT INTO intercept_pending SELECT * FROM json_populate_recordset(NULL::intercept_pending,$1::json)`, string(encoded))
return err
}
func liveTaskExists(tx *sql.Tx, id int64) bool {
var exists bool
_ = tx.QueryRow(`SELECT EXISTS(SELECT 1 FROM tasks WHERE id=$1 AND deleted_at IS NULL)`, id).Scan(&exists)
return exists
}
func nilIfEmptyString(value string) any {
if strings.TrimSpace(value) == "" {
return nil
}
return value
}
// CompleteTaskArchiveRestore removes compact metadata after every external
// component has been verified and the package has been consumed.
func (d *DB) CompleteTaskArchiveRestore(archiveID int64) error {
res, err := d.Exec(`DELETE FROM task_archives archive USING tasks task
WHERE archive.id=$1 AND archive.task_id=task.id AND task.deleted_at IS NULL AND task.archived_at IS NULL`, archiveID)
if err != nil {
return err
}
n, _ := res.RowsAffected()
if n != 1 {
return ErrTaskArchiveState
}
return nil
}
// DeleteTaskArchiveStub permanently removes the cold task after its package has
// been staged for deletion. Dependency protection is rechecked transactionally.
func (d *DB) DeleteTaskArchiveStub(archiveID int64) error {
tx, err := d.Begin()
if err != nil {
return err
}
defer tx.Rollback() //nolint:errcheck
if err := coordinateWithSchemaMigration(tx); err != nil {
return err
}
var taskID, expID int64
var state string
if err := tx.QueryRow(`SELECT archive.task_id,task.exploration_id,archive.state
FROM task_archives archive JOIN tasks task ON task.id=archive.task_id
WHERE archive.id=$1 FOR UPDATE OF archive,task`, archiveID).Scan(&taskID, &expID, &state); err != nil {
return err
}
if state != Deleting {
return ErrTaskArchiveState
}
var dependent int64
err = tx.QueryRow(`SELECT task_id FROM task_archives WHERE id<>$1 AND $2=ANY(source_task_ids) LIMIT 1`, archiveID, taskID).Scan(&dependent)
if err == nil {
return fmt.Errorf("%w: task %d", ErrTaskArchiveDeleteBlocked, dependent)
}
if !errors.Is(err, sql.ErrNoRows) {
return err
}
if _, err := tx.Exec(`DELETE FROM tasks WHERE id=$1 AND archived_at IS NOT NULL`, taskID); err != nil {
return err
}
if _, err := tx.Exec(`DELETE FROM explorations WHERE id=$1`, expID); err != nil {
return err
}
return tx.Commit()
}
// ArchivedAggregateStats returns compact summaries used by global dashboards so
// cold data does not disappear from historical totals.
func (d *DB) ArchivedAggregateStats() ([]json.RawMessage, error) {
rows, err := d.Query(`SELECT archive.aggregate_stats
FROM task_archives archive
JOIN tasks task ON task.id=archive.task_id
WHERE task.archived_at IS NOT NULL`)
if err != nil {
return nil, err
}
defer rows.Close()
var out []json.RawMessage
for rows.Next() {
var raw json.RawMessage
if err := rows.Scan(&raw); err != nil {
return nil, err
}
out = append(out, raw)
}
return out, rows.Err()
}
+468
View File
@@ -0,0 +1,468 @@
package db
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"strings"
"testing"
"time"
)
type archiveJSONTestRows struct {
values [][]byte
next int
}
func (r *archiveJSONTestRows) Next() bool {
if r.next >= len(r.values) {
return false
}
r.next++
return true
}
func (r *archiveJSONTestRows) Scan(dest ...any) error {
if len(dest) != 1 || r.next == 0 || r.next > len(r.values) {
return fmt.Errorf("invalid archive test row scan")
}
target, ok := dest[0].(*[]byte)
if !ok {
return fmt.Errorf("archive test row destination is %T", dest[0])
}
*target = append((*target)[:0], r.values[r.next-1]...)
return nil
}
func (r *archiveJSONTestRows) Err() error { return nil }
func TestQueryArchiveRowsStreamsJSON(t *testing.T) {
payload := strings.Repeat("large request/response payload ", 64*1024)
values := make([][]byte, 2)
var err error
values[0], err = json.Marshal(map[string]any{"id": int64(1), "body": payload})
if err != nil {
t.Fatal(err)
}
values[1], err = json.Marshal(map[string]any{"id": int64(2), "body": "quoted: \"value\"\nline"})
if err != nil {
t.Fatal(err)
}
raw, count, err := encodeArchiveRows(&archiveJSONTestRows{values: values})
if err != nil {
t.Fatal(err)
}
if count != 2 {
t.Fatalf("row count=%d, want 2", count)
}
var rows []struct {
ID int64 `json:"id"`
Body string `json:"body"`
}
if err := json.Unmarshal(raw, &rows); err != nil {
t.Fatal(err)
}
if len(rows) != 2 {
t.Fatalf("decoded row count=%d, want 2", len(rows))
}
if rows[0].ID != 1 || rows[0].Body != payload || rows[1].ID != 2 || rows[1].Body != "quoted: \"value\"\nline" {
t.Fatal("streamed rows were reordered or truncated")
}
}
func TestWriteArchiveRowsProducesJSONSequence(t *testing.T) {
values := [][]byte{
json.RawMessage(`{"id":1,"body":"first"}`),
json.RawMessage(`{"id":2,"body":"second\\nline"}`),
}
var output bytes.Buffer
count, err := writeArchiveRows(&archiveJSONTestRows{values: values}, &output)
if err != nil {
t.Fatal(err)
}
if count != 2 {
t.Fatalf("streamed row count=%d, want 2", count)
}
decoder := json.NewDecoder(&output)
for wantID := int64(1); wantID <= 2; wantID++ {
var row struct {
ID int64 `json:"id"`
}
if err := decoder.Decode(&row); err != nil {
t.Fatal(err)
}
if row.ID != wantID {
t.Fatalf("streamed row id=%d, want %d", row.ID, wantID)
}
}
}
func TestTaskArchiveFormatCompatibility(t *testing.T) {
for _, version := range []int{TaskArchiveLegacyFormatVersion, TaskArchiveFormatVersion} {
if !IsTaskArchiveFormatSupported(version) {
t.Fatalf("archive format %d should be supported", version)
}
}
for _, version := range []int{0, TaskArchiveFormatVersion + 1} {
if IsTaskArchiveFormatSupported(version) {
t.Fatalf("archive format %d should be rejected", version)
}
}
invalidSnapshots := []*TaskArchiveSnapshot{
{FormatVersion: TaskArchiveLegacyFormatVersion, StreamedTables: map[string]string{"llm_records": TaskArchiveLLMRecordsPath}},
{FormatVersion: TaskArchiveFormatVersion, StreamedTables: map[string]string{"unknown": "database/unknown.ndjson"}},
}
for _, snapshot := range invalidSnapshots {
if _, err := (&DB{}).RestoreTaskArchive(1, snapshot, 0); !errors.Is(err, ErrTaskArchiveFormatMismatch) {
t.Fatalf("invalid streamed table metadata returned %v", err)
}
}
}
func TestTaskArchiveDatabaseRoundTrip(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
if err := d.EnsureLLMRecordsTable(); err != nil {
t.Fatal(err)
}
if err := d.EnsureLLMUsageTable(); err != nil {
t.Fatal(err)
}
task, err := d.CreateTaskWithOptions("archive database roundtrip", "restore exact graph", TaskCreateOptions{Name: "cold task"})
if err != nil {
t.Fatal(err)
}
var companyID, llmProfileID int64
if err := d.QueryRow(`INSERT INTO companies(name,nkey) VALUES($1,$2) RETURNING id`,
fmt.Sprintf("archive-company-%d", task.ID), fmt.Sprintf("archive-company-%d", task.ID)).Scan(&companyID); err != nil {
t.Fatal(err)
}
if err := d.QueryRow(`INSERT INTO llm_profiles(name,format,model) VALUES($1,'openai','archive-model') RETURNING id`,
fmt.Sprintf("archive-chain-%d", task.ID)).Scan(&llmProfileID); err != nil {
t.Fatal(err)
}
exhaustedAt := time.Now().UTC().Add(-time.Hour).Truncate(time.Microsecond)
chainCreatedAt := exhaustedAt.Add(-time.Hour)
if _, err := d.Exec(`UPDATE tasks SET company_id=$2,llm_profile_id=$3,active_llm_profile_id=$3 WHERE id=$1`, task.ID, companyID, llmProfileID); err != nil {
t.Fatal(err)
}
if _, err := d.Exec(`INSERT INTO task_llm_profiles(task_id,profile_id,position,status,last_error,exhausted_at,created_at,updated_at)
VALUES($1,$2,0,'quota_exhausted','balance exhausted',$3,$4,$3)`, task.ID, llmProfileID, exhaustedAt, chainCreatedAt); err != nil {
t.Fatal(err)
}
defer func() {
_, _ = d.Exec(`DELETE FROM llm_usage WHERE task_id=$1 OR exploration_id=$2`, fmt.Sprint(task.ID), task.ExplorationID)
_, _ = d.Exec(`DELETE FROM skill_usage WHERE task_id=$1 OR exploration_id=$2`, task.ID, task.ExplorationID)
_, _ = d.Exec(`DELETE FROM tool_usage WHERE task_id=$1 OR exploration_id=$2`, task.ID, task.ExplorationID)
_, _ = d.Exec(`DELETE FROM task_archives WHERE task_id=$1`, task.ID)
_ = d.DeleteTask(task.ID)
_, _ = d.Exec(`DELETE FROM llm_profiles WHERE id=$1`, llmProfileID)
_, _ = d.Exec(`DELETE FROM companies WHERE id=$1`, companyID)
}()
if err := d.SetPaused(task.ID, true); err != nil {
t.Fatal(err)
}
assetID, err := d.Assets().UpsertRootDomain(UpsertRootDomainReq{Domain: fmt.Sprintf("archive-%d.example", task.ID), TaskID: task.ID})
if err != nil {
t.Fatal(err)
}
store := d.Exploration(task.ExplorationID)
nodeID, err := store.AddNode(KindFact, map[string]any{"summary": "archived fact", "asset_ids": []int64{assetID}}, 1, "confirmed", "worker", nil)
if err != nil {
t.Fatal(err)
}
if err := store.Anchor(nodeID, assetID); err != nil {
t.Fatal(err)
}
profileName := fmt.Sprintf("archive-profile-%d", task.ID)
skillName := fmt.Sprintf("archive-skill-%d", task.ID)
toolName := fmt.Sprintf("archive-tool-%d", task.ID)
vulnclass := fmt.Sprintf("archive-vuln-%d", task.ID)
if err := d.InsertLLMUsage(&LLMUsage{TaskID: fmt.Sprint(task.ID), ExplorationID: task.ExplorationID, Worker: "worker", Model: "test", ProfileName: profileName, InputTokens: 11, OutputTokens: 7}); err != nil {
t.Fatal(err)
}
if err := d.InsertSkillUsage(&SkillUsage{Skill: skillName, AgentKey: "worker", TaskID: task.ID, ExplorationID: task.ExplorationID, Found: true}); err != nil {
t.Fatal(err)
}
if err := d.InsertToolUsage(&ToolUsage{ToolKey: toolName, AgentKey: "worker", TaskID: task.ID, ExplorationID: task.ExplorationID}); err != nil {
t.Fatal(err)
}
if err := d.InsertLLMRecord(&LLMRecord{
TaskID: fmt.Sprint(task.ID), Model: "archive-model", SessionID: "archive-session",
Status: "ok", RawRequest: strings.Repeat("request", 1024), RawResponse: strings.Repeat("response", 1024),
}); err != nil {
t.Fatal(err)
}
if _, err := d.AddFinding(task.ID, nodeID, vulnclass, "archive finding", SeverityCritical, "summary", "evidence", "worker", []int64{assetID}); err != nil {
t.Fatal(err)
}
archive, err := d.QueueTaskArchive(task.ID)
if err != nil {
t.Fatal(err)
}
claimed, err := d.ClaimTaskArchiveJob(t.Context())
if err != nil || claimed == nil || claimed.ID != archive.ID || claimed.State != Archiving {
t.Fatalf("claim = %+v, %v", claimed, err)
}
var llmRecords bytes.Buffer
snapshot, err := d.SnapshotTaskArchiveWithLLMRecords(task.ID, &llmRecords)
if err != nil {
t.Fatal(err)
}
if snapshot.StreamedTables["llm_records"] != TaskArchiveLLMRecordsPath || snapshot.DataCounts["llm_records"] != 1 {
t.Fatalf("unexpected streamed LLM metadata: paths=%v counts=%v", snapshot.StreamedTables, snapshot.DataCounts)
}
if rawRowCount(snapshot.Tables["llm_records"]) != 0 {
t.Fatal("streamed LLM records were also retained in manifest memory")
}
if snapshot.DataCounts["assets"] != 1 || snapshot.DataCounts["exploration_nodes"] < 2 {
t.Fatalf("unexpected snapshot counts: %#v", snapshot.DataCounts)
}
if err := d.CompleteTaskArchive(archive.ID, snapshot, "/tmp/test-task.tar.zst", "abc", 100, 50); err != nil {
t.Fatal(err)
}
if live, err := d.GetTask(task.ID); err != nil || live != nil {
t.Fatalf("archived task must be hidden, got %+v, %v", live, err)
}
var nodes, assets, usage int
if err := d.QueryRow(`SELECT count(*) FROM exploration_nodes WHERE exploration_id=$1`, task.ExplorationID).Scan(&nodes); err != nil {
t.Fatal(err)
}
if err := d.QueryRow(`SELECT count(*) FROM assets WHERE id=$1`, assetID).Scan(&assets); err != nil {
t.Fatal(err)
}
if err := d.QueryRow(`SELECT count(*) FROM llm_usage WHERE task_id=$1`, fmt.Sprint(task.ID)).Scan(&usage); err != nil {
t.Fatal(err)
}
if nodes != 0 || assets != 0 || usage != 0 {
t.Fatalf("hot compaction left nodes=%d assets=%d usage=%d", nodes, assets, usage)
}
ready, err := d.GetTaskArchive(archive.ID)
if err != nil || ready == nil || ready.State != ArchiveReady {
t.Fatalf("ready archive = %+v, %v", ready, err)
}
var stats map[string]any
if err := json.Unmarshal(ready.AggregateStats, &stats); err != nil {
t.Fatal(err)
}
if _, ok := stats["tokens"]; !ok {
t.Fatalf("archive token summary missing: %#v", stats)
}
assertArchiveGlobalStats(t, d, profileName, skillName, toolName, vulnclass)
if _, err := d.Exec(`DELETE FROM companies WHERE id=$1`, companyID); err != nil {
t.Fatal(err)
}
if _, err := d.QueueTaskArchiveRestore(archive.ID); err != nil {
t.Fatal(err)
}
claimed, err = d.ClaimTaskArchiveJob(t.Context())
if err != nil || claimed == nil || claimed.State != Restoring {
t.Fatalf("restore claim = %+v, %v", claimed, err)
}
warnings, err := d.RestoreTaskArchiveWithLLMRecords(
archive.ID, snapshot, int64((10 * time.Minute).Seconds()), bytes.NewReader(llmRecords.Bytes()),
)
if err != nil {
t.Fatal(err)
}
warningFound := false
for _, warning := range warnings {
if strings.Contains(warning, fmt.Sprintf("任务企业 %d 已删除", companyID)) {
warningFound = true
}
}
if !warningFound {
t.Fatalf("missing deleted company warning: %v", warnings)
}
live, err := d.GetTask(task.ID)
if err != nil || live == nil {
t.Fatalf("restored task = %+v, %v", live, err)
}
if live.Name != "cold task" || !live.Paused {
t.Fatalf("restored task metadata mismatch: %+v", live)
}
if err := d.QueryRow(`SELECT count(*) FROM exploration_nodes WHERE exploration_id=$1`, task.ExplorationID).Scan(&nodes); err != nil {
t.Fatal(err)
}
if err := d.QueryRow(`SELECT count(*) FROM assets WHERE id=$1 AND $2=ANY(task_ids)`, assetID, task.ID).Scan(&assets); err != nil {
t.Fatal(err)
}
if err := d.QueryRow(`SELECT count(*) FROM llm_usage WHERE task_id=$1`, fmt.Sprint(task.ID)).Scan(&usage); err != nil {
t.Fatal(err)
}
if nodes < 2 || assets != 1 || usage != 1 {
t.Fatalf("restore incomplete nodes=%d assets=%d usage=%d", nodes, assets, usage)
}
var restoredRawRequest, restoredRawResponse string
if err := d.QueryRow(`SELECT COALESCE(raw_request,''),COALESCE(raw_response,'') FROM llm_records WHERE task_id=$1`, fmt.Sprint(task.ID)).
Scan(&restoredRawRequest, &restoredRawResponse); err != nil {
t.Fatal(err)
}
if restoredRawRequest != strings.Repeat("request", 1024) || restoredRawResponse != strings.Repeat("response", 1024) {
t.Fatal("restored streamed LLM record body was truncated")
}
var restoredCompanyID *int64
var restoredStatus, restoredError string
var restoredExhaustedAt, restoredCreatedAt time.Time
if err := d.QueryRow(`SELECT company_id FROM tasks WHERE id=$1`, task.ID).Scan(&restoredCompanyID); err != nil {
t.Fatal(err)
}
if restoredCompanyID != nil {
t.Fatalf("deleted legacy company restored as %v", *restoredCompanyID)
}
if err := d.QueryRow(`SELECT status,COALESCE(last_error,''),exhausted_at,created_at FROM task_llm_profiles WHERE task_id=$1 AND profile_id=$2`,
task.ID, llmProfileID).Scan(&restoredStatus, &restoredError, &restoredExhaustedAt, &restoredCreatedAt); err != nil {
t.Fatal(err)
}
if restoredStatus != "quota_exhausted" || restoredError != "balance exhausted" ||
!restoredExhaustedAt.Equal(exhaustedAt) || !restoredCreatedAt.Equal(chainCreatedAt) {
t.Fatalf("LLM chain state/time mismatch: status=%s error=%q exhausted=%s created=%s",
restoredStatus, restoredError, restoredExhaustedAt, restoredCreatedAt)
}
assertArchiveGlobalStats(t, d, profileName, skillName, toolName, vulnclass)
if err := d.CompleteTaskArchiveRestore(archive.ID); err != nil {
t.Fatal(err)
}
if item, err := d.GetTaskArchive(archive.ID); err != nil || item != nil {
t.Fatalf("archive metadata must be consumed: %+v, %v", item, err)
}
}
func TestTaskArchiveBlockersIgnoreQueuedDependents(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
source, err := d.CreateTaskWithOptions("archive blocker source", "source", TaskCreateOptions{})
if err != nil {
t.Fatal(err)
}
dependent, err := d.CreateTaskWithOptions("archive blocker dependent", "dependent", TaskCreateOptions{
SourceTaskIDs: []int64{source.ID},
})
if err != nil {
_ = d.DeleteTask(source.ID)
t.Fatal(err)
}
defer func() {
_, _ = d.Exec(`DELETE FROM task_archives WHERE task_id IN ($1,$2)`, source.ID, dependent.ID)
_ = d.DeleteTask(dependent.ID)
_ = d.DeleteTask(source.ID)
}()
blockers, err := d.TaskArchiveBlockers()
if err != nil {
t.Fatal(err)
}
if blockers[source.ID] != dependent.ID {
t.Fatalf("source blocker=%d, want dependent %d", blockers[source.ID], dependent.ID)
}
if err := d.SetPaused(dependent.ID, true); err != nil {
t.Fatal(err)
}
if _, err := d.QueueTaskArchive(dependent.ID); err != nil {
t.Fatal(err)
}
blockers, err = d.TaskArchiveBlockers()
if err != nil {
t.Fatal(err)
}
if blocker := blockers[source.ID]; blocker != 0 {
t.Fatalf("queued dependent must not block source, got %d", blocker)
}
}
func TestRecoverInterruptedArchiveRequiresManualRetry(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
task, err := d.CreateTaskWithOptions("interrupted archive", "must not restart automatically", TaskCreateOptions{})
if err != nil {
t.Fatal(err)
}
defer func() {
_, _ = d.Exec(`DELETE FROM task_archives WHERE task_id=$1`, task.ID)
_ = d.DeleteTask(task.ID)
}()
if err := d.SetPaused(task.ID, true); err != nil {
t.Fatal(err)
}
queued, err := d.QueueTaskArchive(task.ID)
if err != nil {
t.Fatal(err)
}
if _, err := d.Exec(`UPDATE task_archives SET state=$2,phase='snapshot_database',progress=10 WHERE id=$1`, queued.ID, Archiving); err != nil {
t.Fatal(err)
}
if err := d.RecoverTaskArchiveJobs(); err != nil {
t.Fatal(err)
}
recovered, err := d.GetTaskArchive(queued.ID)
if err != nil {
t.Fatal(err)
}
if recovered == nil || recovered.State != ArchiveFailed || recovered.Phase != "interrupted" || recovered.Error == "" {
t.Fatalf("recovered archive=%+v, want explicit-retry failure", recovered)
}
}
func assertArchiveGlobalStats(t *testing.T, d *DB, profileName, skillName, toolName, vulnclass string) {
t.Helper()
profiles, err := d.UsageByProfile()
if err != nil {
t.Fatal(err)
}
matchedProfile := false
for _, profile := range profiles {
if profile.ProfileName != profileName {
continue
}
matchedProfile = true
if profile.Calls != 1 || profile.Tasks != 1 || profile.InputTokens != 11 || profile.OutputTokens != 7 {
t.Fatalf("archive profile aggregate double-counted or missing: %+v", profile)
}
}
if !matchedProfile {
t.Fatalf("archive profile aggregate %q missing", profileName)
}
skills, err := d.SkillStats()
if err != nil {
t.Fatal(err)
}
matchedSkill := false
for _, skill := range skills {
if skill.Skill == skillName {
matchedSkill = skill.Calls == 1 && skill.Tasks == 1
}
}
if !matchedSkill {
t.Fatalf("archive skill aggregate %q missing or double-counted: %+v", skillName, skills)
}
tools, err := d.ToolUsageCounts()
if err != nil {
t.Fatal(err)
}
if tools[toolName] != 1 {
t.Fatalf("archive tool aggregate %q=%d, want 1", toolName, tools[toolName])
}
findings, err := d.FindingStats()
if err != nil {
t.Fatal(err)
}
foundClass := false
for _, item := range findings.VulnClasses {
if item == vulnclass {
foundClass = true
}
}
if !foundClass {
t.Fatalf("archive finding class %q missing", vulnclass)
}
}
+425
View File
@@ -0,0 +1,425 @@
package db
import (
"database/sql"
"errors"
"fmt"
"net"
"strings"
"unicode/utf8"
)
const (
MaxTaskAssetMutationCount = 100
MaxTaskAssetSummaryRunes = 500
defaultTaskAssetSource = "system"
manualTaskScopeSummary = "사용자가 테스트 자산 페이지에서 직접 추가"
)
var (
ErrTaskAssetInvalid = errors.New("invalid task asset association")
ErrTaskAssetTaskNotFound = errors.New("task not found")
ErrTaskAssetAssetNotFound = errors.New("asset not found")
)
// TaskAssetMutation summarizes one attach request. Attached counts newly added
// associations; Existing counts requested assets that were already on the task.
type TaskAssetMutation struct {
Requested int `json:"requested"`
Attached int `json:"attached"`
Existing int `json:"existing"`
}
// TaskAssetScopeMutation summarizes one free-form scope registration. Domain
// and IP entries create or reuse global assets; every entry also becomes an
// idempotent task_scope row.
type TaskAssetScopeMutation struct {
Requested int `json:"requested"`
AssetsLinked int `json:"assets_linked"`
AssetsExisting int `json:"assets_existing"`
ScopesAdded int `json:"scopes_added"`
ScopesExisting int `json:"scopes_existing"`
}
// IntentAsset describes an asset explicitly anchored to a worker intent.
type IntentAsset struct {
IntentID int64 `json:"intent_id"`
AssetID int64 `json:"asset_id"`
Type string `json:"type"`
Label string `json:"label"`
Source string `json:"source"`
SourceSummary string `json:"source_summary"`
SourceNodeID *int64 `json:"source_node_id,omitempty"`
SourceTaskID int64 `json:"source_task_id"`
Inherited bool `json:"inherited"`
}
func normalizeTaskAssetIDs(ids []int64) ([]int64, error) {
if len(ids) == 0 {
return nil, fmt.Errorf("%w: asset_ids is required", ErrTaskAssetInvalid)
}
seen := make(map[int64]struct{}, len(ids))
normalized := make([]int64, 0, len(ids))
for _, id := range ids {
if id <= 0 {
return nil, fmt.Errorf("%w: asset id must be positive", ErrTaskAssetInvalid)
}
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
normalized = append(normalized, id)
if len(normalized) > MaxTaskAssetMutationCount {
return nil, fmt.Errorf("%w: at most %d assets per request", ErrTaskAssetInvalid, MaxTaskAssetMutationCount)
}
}
return normalized, nil
}
func normalizeTaskAssetSource(source, summary string) (string, string, error) {
source = strings.TrimSpace(strings.ToLower(source))
if source == "" {
source = defaultTaskAssetSource
}
summary = strings.TrimSpace(summary)
if utf8.RuneCountInString(summary) > MaxTaskAssetSummaryRunes {
return "", "", fmt.Errorf("%w: source summary exceeds %d characters", ErrTaskAssetInvalid, MaxTaskAssetSummaryRunes)
}
return source, summary, nil
}
// SetTaskAssetSource improves the generic trigger-created provenance for one
// existing task association. It never creates or deletes an asset.
func (s *AssetStore) SetTaskAssetSource(taskID, assetID int64, source, summary string, sourceNodeID *int64) error {
if taskID <= 0 || assetID <= 0 {
return fmt.Errorf("%w: task and asset ids must be positive", ErrTaskAssetInvalid)
}
source, summary, err := normalizeTaskAssetSource(source, summary)
if err != nil {
return err
}
query := `
INSERT INTO task_asset_links(task_id, asset_id, source, source_summary, source_node_id)
SELECT task.id, asset.id, $3, $4, $5
FROM tasks task
JOIN assets asset ON asset.id=$2 AND task.id=ANY(asset.task_ids)
WHERE task.id=$1 AND task.deleted_at IS NULL
ON CONFLICT (task_id, asset_id) DO UPDATE
SET source=EXCLUDED.source,
source_summary=EXCLUDED.source_summary,
source_node_id=COALESCE(EXCLUDED.source_node_id, task_asset_links.source_node_id)`
var result sql.Result
if s.tx != nil {
result, err = s.tx.Exec(query, taskID, assetID, source, summary, sourceNodeID)
} else {
result, err = s.db.Exec(query, taskID, assetID, source, summary, sourceNodeID)
}
if err != nil {
return err
}
rows, err := result.RowsAffected()
if err != nil {
return err
}
if rows == 0 {
return fmt.Errorf("%w: task or asset association does not exist", ErrTaskAssetInvalid)
}
return nil
}
// RegisterTaskAssetScopes accepts the same structured scope rules as enterprise
// assets. The entire request is atomic: invalid input or any storage failure
// leaves both global assets and task scope unchanged.
func (s *AssetStore) RegisterTaskAssetScopes(taskID int64, inputs []ScopeInput) (TaskAssetScopeMutation, error) {
mutation := TaskAssetScopeMutation{Requested: len(inputs)}
if taskID <= 0 {
return mutation, fmt.Errorf("%w: task id must be positive", ErrTaskAssetInvalid)
}
if len(inputs) == 0 {
return mutation, fmt.Errorf("%w: scope is required", ErrTaskAssetInvalid)
}
if err := ValidateCompanyScopeInputBounds(inputs); err != nil {
return mutation, fmt.Errorf("%w: %v", ErrTaskAssetInvalid, err)
}
parsed := make([]ParsedScope, 0, len(inputs))
for index, input := range inputs {
rule, err := ParseScopeInput(input)
if err != nil {
return mutation, fmt.Errorf("%w: 第 %d 条范围无效: %v", ErrTaskAssetInvalid, index+1, err)
}
parsed = append(parsed, rule)
}
if err := validateParsedScopeBounds(parsed); err != nil {
return mutation, fmt.Errorf("%w: %v", ErrTaskAssetInvalid, err)
}
tx, err := s.db.Begin()
if err != nil {
return mutation, err
}
defer tx.Rollback() //nolint:errcheck
if err := lockCompanyScopeMutation(tx); err != nil {
return mutation, err
}
var taskExists bool
if err := tx.QueryRow(`SELECT EXISTS(SELECT 1 FROM tasks WHERE id=$1 AND deleted_at IS NULL)`, taskID).Scan(&taskExists); err != nil {
return mutation, err
}
if !taskExists {
return mutation, ErrTaskAssetTaskNotFound
}
scoped := &AssetStore{db: s.db, company: s.company, tx: tx}
for _, rule := range parsed {
taskScope := TaskScope{
TaskID: taskID,
Source: "manual",
Reason: manualTaskScopeSummary,
}
var assetID int64
switch rule.Kind {
case "domain":
taskScope.Kind = "root_domain"
taskScope.Domain = rule.Domain
var alreadyLinked bool
err := tx.QueryRow(`SELECT id, $2=ANY(task_ids) FROM assets WHERE type='root_domain' AND domain=$1`, rule.Domain, taskID).
Scan(&assetID, &alreadyLinked)
if err != nil && !errors.Is(err, sql.ErrNoRows) {
return mutation, err
}
assetID, err = scoped.UpsertRootDomain(UpsertRootDomainReq{Domain: rule.Domain, TaskID: taskID})
if err != nil {
return mutation, err
}
if alreadyLinked {
mutation.AssetsExisting++
} else {
mutation.AssetsLinked++
}
case "ip":
taskScope.Kind = "ip"
taskScope.Net = rule.Net
ip, _, parseErr := net.ParseCIDR(rule.Net)
if parseErr != nil {
return mutation, fmt.Errorf("%w: 无效 IP: %s", ErrTaskAssetInvalid, rule.Raw)
}
ipValue := ip.String()
var alreadyLinked bool
err := tx.QueryRow(`SELECT id, $2=ANY(task_ids) FROM assets WHERE type='ip' AND ip=$1`, ipValue, taskID).
Scan(&assetID, &alreadyLinked)
if err != nil && !errors.Is(err, sql.ErrNoRows) {
return mutation, err
}
assetID, err = scoped.UpsertIP(UpsertIPReq{IP: ipValue, TaskID: taskID})
if err != nil {
return mutation, err
}
if alreadyLinked {
mutation.AssetsExisting++
} else {
mutation.AssetsLinked++
}
case "cidr":
taskScope.Kind = "cidr"
taskScope.Net = rule.Net
case "icp", "keyword":
taskScope.Kind = rule.Kind
taskScope.Value = rule.Value
default:
return mutation, fmt.Errorf("%w: unsupported scope kind %q", ErrTaskAssetInvalid, rule.Kind)
}
if assetID > 0 {
if err := scoped.SetTaskAssetSource(taskID, assetID, "manual", manualTaskScopeSummary, nil); err != nil {
return mutation, err
}
}
inserted, err := scoped.upsertTaskScopeResult(taskScope)
if err != nil {
return mutation, err
}
if inserted {
mutation.ScopesAdded++
} else {
mutation.ScopesExisting++
}
}
if err := tx.Commit(); err != nil {
return mutation, err
}
return mutation, nil
}
// AttachAssetsToTask associates existing global assets with one live task and
// records an operator-authored source summary. Global asset rows are retained.
func (s *AssetStore) AttachAssetsToTask(taskID int64, assetIDs []int64, sourceSummary string) (TaskAssetMutation, error) {
var mutation TaskAssetMutation
assetIDs, err := normalizeTaskAssetIDs(assetIDs)
if err != nil {
return mutation, err
}
_, sourceSummary, err = normalizeTaskAssetSource("manual", sourceSummary)
if err != nil {
return mutation, err
}
if sourceSummary == "" {
return mutation, fmt.Errorf("%w: source_summary is required", ErrTaskAssetInvalid)
}
mutation.Requested = len(assetIDs)
tx, err := s.db.Begin()
if err != nil {
return mutation, err
}
defer tx.Rollback() //nolint:errcheck
var taskExists bool
if err := tx.QueryRow(`SELECT EXISTS(SELECT 1 FROM tasks WHERE id=$1 AND deleted_at IS NULL)`, taskID).Scan(&taskExists); err != nil {
return mutation, err
}
if !taskExists {
return mutation, ErrTaskAssetTaskNotFound
}
var found, existing int
if err := tx.QueryRow(`
SELECT count(*), count(*) FILTER (WHERE $1=ANY(task_ids))
FROM assets WHERE id=ANY($2::bigint[])`, taskID, assetIDs).Scan(&found, &existing); err != nil {
return mutation, err
}
if found != len(assetIDs) {
return mutation, ErrTaskAssetAssetNotFound
}
// Order matters: this UPDATE fires trg_assets_task_links, which creates the
// link rows with the generic source='system'. The INSERT below must stay
// after it so the operator-authored 'manual' provenance wins; swapping the
// two statements silently degrades every manual attach back to 'system'.
if _, err := tx.Exec(`
UPDATE assets
SET task_ids=CASE WHEN $1=ANY(task_ids) THEN task_ids ELSE array_append(task_ids,$1) END
WHERE id=ANY($2::bigint[])`, taskID, assetIDs); err != nil {
return mutation, err
}
if _, err := tx.Exec(`
INSERT INTO task_asset_links(task_id, asset_id, source, source_summary)
SELECT $1, id, 'manual', $3 FROM assets WHERE id=ANY($2::bigint[])
ON CONFLICT (task_id, asset_id) DO UPDATE
SET source='manual', source_summary=EXCLUDED.source_summary, source_node_id=NULL`,
taskID, assetIDs, sourceSummary); err != nil {
return mutation, err
}
mutation.Existing = existing
mutation.Attached = len(assetIDs) - existing
return mutation, tx.Commit()
}
// DetachAssetFromTask removes only the task association. The global asset and
// exploration anchors remain available for historical blackboard auditing.
func (s *AssetStore) DetachAssetFromTask(taskID, assetID int64) (bool, error) {
if taskID <= 0 || assetID <= 0 {
return false, fmt.Errorf("%w: task and asset ids must be positive", ErrTaskAssetInvalid)
}
var detachedID int64
err := s.db.QueryRow(`
UPDATE assets SET task_ids=array_remove(task_ids,$1)
WHERE id=$2 AND $1=ANY(task_ids)
RETURNING id`, taskID, assetID).Scan(&detachedID)
if err == sql.ErrNoRows {
return false, nil
}
if err != nil {
return false, err
}
return detachedID == assetID, nil
}
func (s *AssetStore) hydrateTaskAssetSources(taskID int64, assets []*Asset) error {
if len(assets) == 0 {
return nil
}
ids := make([]int64, 0, len(assets))
byID := make(map[int64]*Asset, len(assets))
for _, asset := range assets {
ids = append(ids, asset.ID)
byID[asset.ID] = asset
}
rows, err := s.db.Query(`
SELECT asset_id, source, source_summary, source_node_id
FROM task_asset_links
WHERE task_id=$1 AND asset_id=ANY($2::bigint[])`, taskID, ids)
if err != nil {
return err
}
defer rows.Close()
for rows.Next() {
var assetID int64
var source, summary string
var sourceNodeID sql.NullInt64
if err := rows.Scan(&assetID, &source, &summary, &sourceNodeID); err != nil {
return err
}
if asset := byID[assetID]; asset != nil {
asset.TaskSource = source
asset.TaskSourceSummary = summary
if sourceNodeID.Valid {
id := sourceNodeID.Int64
asset.TaskSourceNodeID = &id
}
}
}
return rows.Err()
}
// IntentAssets returns all local worker targets plus immutable targets from the
// task's direct sources. Inherited non-terminal intents remain hidden, matching
// the existing source-aware session contract.
func (s *AssetStore) IntentAssets(taskID int64) ([]IntentAsset, error) {
rows, err := s.db.Query(`
WITH context AS (
SELECT task.id AS task_id, task.exploration_id, false AS inherited
FROM tasks task
WHERE task.id=$1 AND task.deleted_at IS NULL
UNION ALL
SELECT source.id, source.exploration_id, true
FROM task_relations relation
JOIN tasks source ON source.id=relation.source_task_id AND source.deleted_at IS NULL
WHERE relation.task_id=$1
)
SELECT intent.id, asset.id, asset.type,
CASE asset.type
WHEN 'root_domain' THEN COALESCE(asset.domain,'')
WHEN 'subdomain' THEN COALESCE(asset.domain,'')
WHEN 'ip' THEN COALESCE(asset.ip,'')
WHEN 'app' THEN COALESCE(asset.app_name,'')
WHEN 'service' THEN COALESCE(NULLIF(asset.url,''), NULLIF(concat_ws(':', COALESCE(NULLIF(asset.domain,''), NULLIF(asset.ip,'')), asset.port::text),''), NULLIF(asset.service_name,''), '#' || asset.id::text)
WHEN 'endpoint' THEN COALESCE(NULLIF(asset.url,''), '#' || asset.id::text)
ELSE '#' || asset.id::text
END,
COALESCE(link.source,'anchor'),
COALESCE(NULLIF(link.source_summary,''), '의도가 앵커로 연결한 자산'),
link.source_node_id, context.task_id, context.inherited
FROM context
JOIN exploration_nodes intent ON intent.exploration_id=context.exploration_id AND intent.kind='intent'
JOIN exploration_anchors anchor ON anchor.node_id=intent.id
JOIN assets asset ON asset.id=anchor.asset_id
LEFT JOIN task_asset_links link ON link.task_id=context.task_id AND link.asset_id=asset.id
WHERE NOT context.inherited OR intent.state IN ('done','blocked','exhausted','stopped')
ORDER BY context.inherited, intent.id DESC, asset.id`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
out := []IntentAsset{}
for rows.Next() {
var asset IntentAsset
var sourceNodeID sql.NullInt64
if err := rows.Scan(&asset.IntentID, &asset.AssetID, &asset.Type, &asset.Label,
&asset.Source, &asset.SourceSummary, &sourceNodeID, &asset.SourceTaskID, &asset.Inherited); err != nil {
return nil, err
}
if sourceNodeID.Valid {
id := sourceNodeID.Int64
asset.SourceNodeID = &id
}
out = append(out, asset)
}
return out, rows.Err()
}
+260
View File
@@ -0,0 +1,260 @@
package db
import (
"net/url"
"sort"
"strconv"
"strings"
)
// directTaskContextCTE is the current task plus exactly its explicitly related
// source tasks. It is deliberately non-recursive.
const directTaskContextCTE = `
context_tasks AS (
SELECT t.id AS task_id, t.exploration_id
FROM tasks t
WHERE t.id=$1 AND t.deleted_at IS NULL
UNION ALL
SELECT source.id, source.exploration_id
FROM task_relations relation
JOIN tasks source ON source.id=relation.source_task_id AND source.deleted_at IS NULL
WHERE relation.task_id=$1
)`
const contextCoverageCTE = directTaskContextCTE + `,
target AS (
SELECT DISTINCT a.id, a.type,
COALESCE(a.url, a.domain, a.ip, a.app_name, a.root_domain, '') AS label
FROM assets a
JOIN task_scope ts ON (
(ts.kind='company' AND a.company_id = ts.company_id)
OR (ts.kind='root_domain' AND a.root_domain = ts.domain)
OR (ts.kind='subdomain' AND a.domain = ts.domain)
OR (ts.kind IN ('ip','cidr') AND ts.net >>= try_inet(a.ip))
OR (ts.kind='icp' AND (
lower(regexp_replace(COALESCE(a.icp,''), '[[:space:]]+', '', 'g')) = ts.value
OR lower(regexp_replace(COALESCE(a.app_icp,''), '[[:space:]]+', '', 'g')) = ts.value
))
)
JOIN context_tasks ctx ON ctx.task_id=ts.task_id
UNION
SELECT a.id, a.type,
COALESCE(a.url, a.domain, a.ip, a.app_name, a.root_domain, '') AS label
FROM assets a
JOIN exploration_anchors ea ON ea.asset_id=a.id
JOIN exploration_nodes en ON en.id=ea.node_id
JOIN context_tasks ctx ON ctx.exploration_id=en.exploration_id
),
tested AS (
SELECT DISTINCT ea.asset_id
FROM exploration_anchors ea
JOIN exploration_nodes en ON en.id=ea.node_id AND en.kind='fact'
JOIN context_tasks ctx ON ctx.exploration_id=en.exploration_id
)`
// scopeTargetCTE selects every asset that BELONGS to the current task's (and its
// direct source tasks') declared scope — membership, not literal value: a
// root_domain scope pulls in every subdomain / service / endpoint whose own
// root_domain column equals it; an ip/cidr scope pulls in assets whose ip OR
// IP-literal host falls inside the net. $1 is the task id. Unlike contextCoverageCTE
// it carries neither the fact-anchor union nor the tested set — it is pure "in
// declared scope", independent of what has already been touched. Used by the
// agent's list_assets so a query returns the task's relevant assets, not the
// whole shared库.
const scopeTargetCTE = `
context_tasks AS (
SELECT t.id AS task_id
FROM tasks t
WHERE t.id=$1 AND t.deleted_at IS NULL
UNION ALL
SELECT source.id
FROM task_relations relation
JOIN tasks source ON source.id=relation.source_task_id AND source.deleted_at IS NULL
WHERE relation.task_id=$1
),
target AS (
SELECT DISTINCT a.id
FROM assets a
JOIN task_scope ts ON (
(ts.kind='company' AND a.company_id = ts.company_id)
OR (ts.kind='root_domain' AND a.root_domain = ts.domain)
OR (ts.kind='subdomain' AND a.domain = ts.domain)
OR (ts.kind IN ('ip','cidr') AND (ts.net >>= try_inet(a.ip) OR ts.net >>= try_inet(a.domain)))
OR (ts.kind='icp' AND (
lower(regexp_replace(COALESCE(a.icp,''), '[[:space:]]+', '', 'g')) = ts.value
OR lower(regexp_replace(COALESCE(a.app_icp,''), '[[:space:]]+', '', 'g')) = ts.value
))
)
JOIN context_tasks ctx ON ctx.task_id=ts.task_id
)`
// ListTaskScopeWithSources returns the current task's scope followed by the
// scopes of its direct source tasks. TaskScope.TaskID preserves provenance.
func (s *AssetStore) ListTaskScopeWithSources(taskID int64) ([]TaskScope, error) {
rows, err := s.db.Query(`WITH `+directTaskContextCTE+`
SELECT ts.id, ts.task_id, ts.kind, COALESCE(ts.company_id,0), COALESCE(c.name,''), COALESCE(ts.domain,''),
COALESCE(ts.net::text,''), COALESCE(ts.value,''), ts.source, COALESCE(ts.reason,'')
FROM task_scope ts
JOIN context_tasks ctx ON ctx.task_id=ts.task_id
LEFT JOIN companies c ON c.id=ts.company_id
ORDER BY CASE WHEN ts.task_id=$1 THEN 0 ELSE 1 END, ts.id`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
out := []TaskScope{}
for rows.Next() {
var scope TaskScope
var companyID int64
if err := rows.Scan(&scope.ID, &scope.TaskID, &scope.Kind, &companyID, &scope.CompanyName, &scope.Domain, &scope.Net, &scope.Value, &scope.Source, &scope.Reason); err != nil {
return nil, err
}
if companyID > 0 {
scope.CompanyID = &companyID
}
out = append(out, scope)
}
return out, rows.Err()
}
// TaskCoverageWithSources computes one coverage view over the union of the
// current task and its direct sources: source scopes and anchored assets extend
// the denominator, while fact anchors count as tested. No row is copied.
func (s *AssetStore) TaskCoverageWithSources(taskID int64) (*Coverage, error) {
cov := &Coverage{ByType: []CoverageByType{}}
_ = s.db.QueryRow(`WITH `+directTaskContextCTE+`
SELECT count(*) FROM task_scope ts JOIN context_tasks ctx ON ctx.task_id=ts.task_id`, taskID).Scan(&cov.ScopeRows)
rows, err := s.db.Query(`WITH `+contextCoverageCTE+`
SELECT target.type, count(*) AS total,
count(*) FILTER (WHERE target.id IN (SELECT asset_id FROM tested)) AS tested
FROM target GROUP BY target.type ORDER BY target.type`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
for rows.Next() {
var byType CoverageByType
if err := rows.Scan(&byType.Type, &byType.Total, &byType.Tested); err != nil {
return nil, err
}
cov.ByType = append(cov.ByType, byType)
cov.Denominator += byType.Total
cov.Tested += byType.Tested
}
if err := rows.Err(); err != nil {
return nil, err
}
if cov.Denominator > 0 {
pct := float64(cov.Tested) / float64(cov.Denominator)
cov.Pct = &pct
}
return cov, nil
}
// ListUntestedAssetsWithSources is the direct-source-aware backlog query used
// by inherited tasks. Source fact anchors remove assets from the backlog.
func (s *AssetStore) ListUntestedAssetsWithSources(taskID int64, typ string, limit, offset int) ([]CoverageAsset, int, error) {
if limit <= 0 {
limit = 10
}
if offset < 0 {
offset = 0
}
typeFilter := ""
args := []any{taskID}
if typ != "" {
typeFilter = " AND target.type = $2"
args = append(args, typ)
}
var total int
if err := s.db.QueryRow(`WITH `+contextCoverageCTE+`
SELECT count(*) FROM target WHERE target.id NOT IN (SELECT asset_id FROM tested)`+typeFilter, args...).Scan(&total); err != nil {
return nil, 0, err
}
pageArgs := append(append([]any{}, args...), limit, offset)
limitPosition := strconv.Itoa(len(args) + 1)
offsetPosition := strconv.Itoa(len(args) + 2)
rows, err := s.db.Query(`WITH `+contextCoverageCTE+`
SELECT target.id, target.type, target.label FROM target
WHERE target.id NOT IN (SELECT asset_id FROM tested)`+typeFilter+`
ORDER BY target.id LIMIT $`+limitPosition+` OFFSET $`+offsetPosition, pageArgs...)
if err != nil {
return nil, total, err
}
defer rows.Close()
out := []CoverageAsset{}
for rows.Next() {
var asset CoverageAsset
if err := rows.Scan(&asset.ID, &asset.Type, &asset.Label); err != nil {
return nil, total, err
}
out = append(out, asset)
}
return out, total, rows.Err()
}
// HostsByTaskWithSources resolves exact HTTP host candidates from assets that
// are attached to, anchored by, or in scope for the current task or a direct
// source. Traffic remains global and is not copied. This read helper must not be
// used for destructive task cleanup; HostsByTask intentionally retains that
// narrower, task-owned behavior.
func (s *AssetStore) HostsByTaskWithSources(taskID int64) ([]string, error) {
rows, err := s.db.Query(`WITH `+directTaskContextCTE+`,
context_assets AS (
SELECT DISTINCT a.id
FROM assets a
WHERE EXISTS (SELECT 1 FROM context_tasks ctx WHERE ctx.task_id=ANY(a.task_ids))
UNION
SELECT ea.asset_id
FROM exploration_anchors ea
JOIN exploration_nodes en ON en.id=ea.node_id
JOIN context_tasks ctx ON ctx.exploration_id=en.exploration_id
UNION
SELECT DISTINCT a.id
FROM assets a
JOIN task_scope ts ON (
(ts.kind='company' AND a.company_id=ts.company_id)
OR (ts.kind='root_domain' AND a.root_domain=ts.domain)
OR (ts.kind='subdomain' AND a.domain=ts.domain)
OR (ts.kind IN ('ip','cidr') AND ts.net >>= try_inet(a.ip))
OR (ts.kind='icp' AND (
lower(regexp_replace(COALESCE(a.icp,''), '[[:space:]]+', '', 'g')) = ts.value
OR lower(regexp_replace(COALESCE(a.app_icp,''), '[[:space:]]+', '', 'g')) = ts.value
))
)
JOIN context_tasks ctx ON ctx.task_id=ts.task_id
)
SELECT COALESCE(a.domain,''), COALESCE(a.ip,''), COALESCE(a.url,'')
FROM assets a JOIN context_assets ctx ON ctx.id=a.id`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
hosts := map[string]struct{}{}
add := func(host string) {
host = strings.TrimSpace(strings.ToLower(host))
if host != "" {
hosts[host] = struct{}{}
}
}
for rows.Next() {
var domain, ip, rawURL string
if err := rows.Scan(&domain, &ip, &rawURL); err != nil {
return nil, err
}
add(domain)
add(ip)
if parsed, err := url.Parse(rawURL); err == nil {
add(parsed.Hostname())
}
}
if err := rows.Err(); err != nil {
return nil, err
}
out := make([]string, 0, len(hosts))
for host := range hosts {
out = append(out, host)
}
sort.Strings(out)
return out, nil
}
+30
View File
@@ -0,0 +1,30 @@
package db
import (
"testing"
"unicode"
)
// TestManualTaskScopeSummaryLocalized pins the manual task-scope provenance label.
// It is stored on task_asset_links.source_summary (AddTaskScope / SetTaskAssetSource)
// and rendered verbatim on the task detail sessions/assets tabs, so reverting it to
// Chinese would surface mixed-language provenance labels in the same panel. The
// existing task_assets_test.go already asserts the stored value equals this constant
// via the symbol, so pinning the constant here protects the user-facing text too.
func TestManualTaskScopeSummaryLocalized(t *testing.T) {
if manualTaskScopeSummary == "" {
t.Fatal("manualTaskScopeSummary: 빈 문자열")
}
hasHangul := false
for _, r := range manualTaskScopeSummary {
if unicode.Is(unicode.Han, r) {
t.Fatalf("manualTaskScopeSummary: 중국어 한자가 남아 있습니다: %q", manualTaskScopeSummary)
}
if unicode.Is(unicode.Hangul, r) {
hasHangul = true
}
}
if !hasHangul {
t.Fatalf("manualTaskScopeSummary: 한글이 없습니다: %q", manualTaskScopeSummary)
}
}
+188
View File
@@ -0,0 +1,188 @@
package db
import (
"errors"
"fmt"
"testing"
"time"
)
func TestRegisterTaskAssetScopesCreatesAssetsAndPersistsTextScope(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
task, err := d.CreateTask("manual scope registration", "goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(task.ID) })
suffix := time.Now().UnixNano()
domain := fmt.Sprintf("manual-scope-%d.example.test", suffix)
ip := fmt.Sprintf("203.0.%d.%d", (suffix/250)%250, suffix%250+1)
inputs := []ScopeInput{
{Kind: "domain", Value: domain},
{Kind: "ip", Value: ip},
{Kind: "cidr", Value: "198.51.100.0/24"},
{Kind: "icp", Value: " 京 ICP 备 12345678 号-1 "},
{Kind: "keyword", Value: " Acme Security "},
}
first, err := d.Assets().RegisterTaskAssetScopes(task.ID, inputs)
if err != nil {
t.Fatal(err)
}
if first.Requested != 5 || first.AssetsLinked != 2 || first.AssetsExisting != 0 || first.ScopesAdded != 5 || first.ScopesExisting != 0 {
t.Fatalf("unexpected first mutation: %+v", first)
}
assets, err := d.Assets().QueryByTask(task.ID, "", 10, 0)
if err != nil || len(assets) != 2 {
t.Fatalf("task assets=%+v err=%v", assets, err)
}
assetIDs := make([]int64, 0, len(assets))
for _, asset := range assets {
assetIDs = append(assetIDs, asset.ID)
if asset.TaskSource != "manual" || asset.TaskSourceSummary != manualTaskScopeSummary {
t.Fatalf("unexpected task asset provenance: %+v", asset)
}
}
t.Cleanup(func() { _, _ = d.Assets().DeleteByIDs(assetIDs) })
scopes, err := d.Assets().ListTaskScope(task.ID)
if err != nil || len(scopes) != 5 {
t.Fatalf("task scope=%+v err=%v", scopes, err)
}
values := map[string]string{}
for _, scope := range scopes {
values[scope.Kind] = scope.Value
}
if values["icp"] != NormalizeICP(inputs[3].Value) || values["keyword"] != normalizeKeyword(inputs[4].Value) {
t.Fatalf("text scope not normalized: %+v", values)
}
second, err := d.Assets().RegisterTaskAssetScopes(task.ID, inputs)
if err != nil {
t.Fatal(err)
}
if second.AssetsExisting != 2 || second.AssetsLinked != 0 || second.ScopesExisting != 5 || second.ScopesAdded != 0 {
t.Fatalf("unexpected idempotent mutation: %+v", second)
}
rollbackDomain := fmt.Sprintf("rollback-%d.example.test", suffix)
_, err = d.Assets().RegisterTaskAssetScopes(task.ID, []ScopeInput{
{Kind: "domain", Value: rollbackDomain},
{Kind: "cidr", Value: "10.0.0.0/8"},
})
if !errors.Is(err, ErrTaskAssetInvalid) {
t.Fatalf("invalid CIDR error=%v", err)
}
var rollbackAssets int
if err := d.QueryRow(`SELECT count(*) FROM assets WHERE type='root_domain' AND domain=$1`, rollbackDomain).Scan(&rollbackAssets); err != nil {
t.Fatal(err)
}
if rollbackAssets != 0 {
t.Fatalf("invalid request created %d assets", rollbackAssets)
}
}
func TestTaskAssetAttachDetachPreservesGlobalAssetAndAnchors(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
task, err := d.CreateTask("asset editing", "goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(task.ID) })
assetID, err := d.Assets().UpsertRootDomain(UpsertRootDomainReq{Domain: fmt.Sprintf("manual-%d.example.test", time.Now().UnixNano())})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _, _ = d.Assets().DeleteByIDs([]int64{assetID}) })
mutation, err := d.Assets().AttachAssetsToTask(task.ID, []int64{assetID, assetID}, "授权资产清单第 3 项")
if err != nil {
t.Fatal(err)
}
if mutation.Requested != 1 || mutation.Attached != 1 || mutation.Existing != 0 {
t.Fatalf("unexpected first mutation: %+v", mutation)
}
assets, err := d.Assets().QueryByTask(task.ID, "", 10, 0)
if err != nil || len(assets) != 1 {
t.Fatalf("task assets=%+v err=%v", assets, err)
}
if assets[0].TaskSource != "manual" || assets[0].TaskSourceSummary != "授权资产清单第 3 项" {
t.Fatalf("unexpected provenance: %+v", assets[0])
}
intentID, err := d.Exploration(task.ExplorationID).AddIntent(map[string]any{"summary": "test asset"}, 5, []int64{assetID}, "human")
if err != nil {
t.Fatal(err)
}
detached, err := d.Assets().DetachAssetFromTask(task.ID, assetID)
if err != nil || !detached {
t.Fatalf("detach=%v err=%v", detached, err)
}
if assets, err := d.Assets().QueryByTask(task.ID, "", 10, 0); err != nil || len(assets) != 0 {
t.Fatalf("detached task assets=%+v err=%v", assets, err)
}
var global, anchors, links int
_ = d.QueryRow(`SELECT count(*) FROM assets WHERE id=$1`, assetID).Scan(&global)
_ = d.QueryRow(`SELECT count(*) FROM exploration_anchors WHERE node_id=$1 AND asset_id=$2`, intentID, assetID).Scan(&anchors)
_ = d.QueryRow(`SELECT count(*) FROM task_asset_links WHERE task_id=$1 AND asset_id=$2`, task.ID, assetID).Scan(&links)
if global != 1 || anchors != 1 || links != 0 {
t.Fatalf("global=%d anchors=%d links=%d", global, anchors, links)
}
}
func TestIntentAssetsIncludesDirectSourceProvenance(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
source, err := d.CreateTask("source assets", "goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
current, err := d.CreateTaskWithOptions("current assets", "goal", TaskCreateOptions{SourceTaskIDs: []int64{source.ID}})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(current.ID); _ = d.DeleteTask(source.ID) })
assetID, err := d.Assets().UpsertHTTPService(UpsertHTTPServiceReq{
URL: fmt.Sprintf("https://intent-%d.example.test", time.Now().UnixNano()), TaskID: source.ID,
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _, _ = d.Assets().DeleteByIDs([]int64{assetID}) })
intentID, err := d.Exploration(source.ExplorationID).AddIntent(map[string]any{"summary": "source worker"}, 5, []int64{assetID}, "planner")
if err != nil {
t.Fatal(err)
}
if err := d.Exploration(source.ExplorationID).SetNodeState(intentID, "done"); err != nil {
t.Fatal(err)
}
nodeID := intentID
if err := d.Assets().SetTaskAssetSource(source.ID, assetID, "agent", "Worker 通过 insert_assets 登记", &nodeID); err != nil {
t.Fatal(err)
}
assets, err := d.Assets().IntentAssets(current.ID)
if err != nil {
t.Fatal(err)
}
if len(assets) != 1 || assets[0].IntentID != intentID || assets[0].SourceTaskID != source.ID || !assets[0].Inherited {
t.Fatalf("unexpected intent assets: %+v", assets)
}
if assets[0].Source != "agent" || assets[0].SourceSummary == "" || assets[0].SourceNodeID == nil || *assets[0].SourceNodeID != intentID {
t.Fatalf("unexpected intent provenance: %+v", assets[0])
}
}
+282
View File
@@ -0,0 +1,282 @@
package db
import (
"database/sql"
"errors"
"fmt"
"strings"
"time"
"unicode/utf8"
"github.com/jackc/pgx/v5/pgconn"
)
const MaxTaskCategoryNameRunes = 80
// MaxTaskCategoryBatchSize bounds one batch move so a single request cannot lock
// an unbounded number of task rows.
const MaxTaskCategoryBatchSize = 100
var (
ErrTaskCategoryInvalid = errors.New("invalid task category")
ErrTaskCategoryNameConflict = errors.New("task category name already exists")
ErrTaskCategoryNotFound = errors.New("task category not found")
ErrTaskCategoryTaskNotFound = errors.New("task not found")
)
// TaskCategory is a globally reusable task grouping label.
type TaskCategory struct {
ID int64 `json:"id"`
Name string `json:"name"`
NKey string `json:"-"`
TaskCount int `json:"task_count"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
const taskCategoryCols = `category.id, category.name, category.nkey,
count(task.id) FILTER (WHERE task.deleted_at IS NULL),
category.created_at, category.updated_at`
func scanTaskCategory(row interface{ Scan(...any) error }) (TaskCategory, error) {
var category TaskCategory
err := row.Scan(&category.ID, &category.Name, &category.NKey, &category.TaskCount, &category.CreatedAt, &category.UpdatedAt)
return category, err
}
func normalizeTaskCategoryName(name string) (string, string, error) {
name = strings.Join(strings.Fields(name), " ")
if name == "" {
return "", "", fmt.Errorf("%w: name is required", ErrTaskCategoryInvalid)
}
if utf8.RuneCountInString(name) > MaxTaskCategoryNameRunes {
return "", "", fmt.Errorf("%w: name exceeds %d characters", ErrTaskCategoryInvalid, MaxTaskCategoryNameRunes)
}
return name, strings.ToLower(name), nil
}
func taskCategoryUniqueViolation(err error) bool {
var pgErr *pgconn.PgError
return errors.As(err, &pgErr) && pgErr.Code == "23505"
}
func (d *DB) CreateTaskCategory(name string) (*TaskCategory, error) {
name, nkey, err := normalizeTaskCategoryName(name)
if err != nil {
return nil, err
}
category, err := scanTaskCategory(d.QueryRow(`
WITH inserted AS (
INSERT INTO task_categories(name, nkey)
VALUES ($1,$2)
ON CONFLICT (nkey) DO NOTHING
RETURNING *
)
SELECT inserted.id, inserted.name, inserted.nkey, 0, inserted.created_at, inserted.updated_at
FROM inserted`, name, nkey))
if err == sql.ErrNoRows {
return nil, ErrTaskCategoryNameConflict
}
if err != nil {
return nil, err
}
return &category, nil
}
func (d *DB) ListTaskCategories() ([]*TaskCategory, error) {
rows, err := d.Query(`
SELECT ` + taskCategoryCols + `
FROM task_categories category
LEFT JOIN tasks task ON task.category_id=category.id
GROUP BY category.id
ORDER BY category.name, category.id`)
if err != nil {
return nil, err
}
defer rows.Close()
categories := []*TaskCategory{}
for rows.Next() {
category, err := scanTaskCategory(rows)
if err != nil {
return nil, err
}
categories = append(categories, &category)
}
return categories, rows.Err()
}
func (d *DB) GetTaskCategory(id int64) (*TaskCategory, error) {
category, err := scanTaskCategory(d.QueryRow(`
SELECT `+taskCategoryCols+`
FROM task_categories category
LEFT JOIN tasks task ON task.category_id=category.id
WHERE category.id=$1
GROUP BY category.id`, id))
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
return &category, nil
}
func (d *DB) RenameTaskCategory(id int64, name string) (*TaskCategory, error) {
name, nkey, err := normalizeTaskCategoryName(name)
if err != nil {
return nil, err
}
category, err := scanTaskCategory(d.QueryRow(`
WITH updated AS (
UPDATE task_categories SET name=$2, nkey=$3 WHERE id=$1 RETURNING *
)
SELECT updated.id, updated.name, updated.nkey,
(SELECT count(*) FROM tasks WHERE category_id=updated.id AND deleted_at IS NULL),
updated.created_at, updated.updated_at
FROM updated`, id, name, nkey))
if err == sql.ErrNoRows {
return nil, ErrTaskCategoryNotFound
}
if taskCategoryUniqueViolation(err) {
return nil, ErrTaskCategoryNameConflict
}
if err != nil {
return nil, err
}
return &category, nil
}
// DeleteTaskCategory moves affected tasks to the uncategorized bucket through
// the tasks.category_id ON DELETE SET NULL foreign key.
func (d *DB) DeleteTaskCategory(id int64) (bool, error) {
result, err := d.Exec(`DELETE FROM task_categories WHERE id=$1`, id)
if err != nil {
return false, err
}
rows, err := result.RowsAffected()
return rows > 0, err
}
// SetTaskCategory updates one live task. A nil category means uncategorized.
func (d *DB) SetTaskCategory(taskID int64, categoryID *int64) (*TaskCategory, error) {
if categoryID == nil {
result, err := d.Exec(`UPDATE tasks SET category_id=NULL WHERE id=$1 AND deleted_at IS NULL`, taskID)
if err != nil {
return nil, err
}
rows, err := result.RowsAffected()
if err != nil {
return nil, err
}
if rows == 0 {
return nil, ErrTaskCategoryTaskNotFound
}
return nil, nil
}
if *categoryID <= 0 {
return nil, fmt.Errorf("%w: category id must be positive", ErrTaskCategoryInvalid)
}
category, err := scanTaskCategory(d.QueryRow(`
WITH selected AS (
SELECT * FROM task_categories WHERE id=$2
), updated AS (
UPDATE tasks SET category_id=$2
WHERE id=$1 AND deleted_at IS NULL AND EXISTS (SELECT 1 FROM selected)
RETURNING id
)
SELECT selected.id, selected.name, selected.nkey,
(SELECT count(*) FROM tasks WHERE category_id=selected.id AND deleted_at IS NULL),
selected.created_at, selected.updated_at
FROM selected, updated`, taskID, *categoryID))
if err == sql.ErrNoRows {
var taskExists bool
if checkErr := d.QueryRow(`SELECT EXISTS(SELECT 1 FROM tasks WHERE id=$1 AND deleted_at IS NULL)`, taskID).Scan(&taskExists); checkErr != nil {
return nil, checkErr
}
if !taskExists {
return nil, ErrTaskCategoryTaskNotFound
}
return nil, ErrTaskCategoryNotFound
}
if err != nil {
return nil, err
}
return &category, nil
}
// SetTasksCategory moves several tasks into one category (nil = uncategorized)
// inside a single transaction, so a half-applied batch is never observable.
// It returns the ids that were actually updated — ids missing from that slice
// were deleted between selection and submit — plus the refreshed category row
// whose task_count already reflects this move.
func (d *DB) SetTasksCategory(taskIDs []int64, categoryID *int64) ([]int64, *TaskCategory, error) {
if len(taskIDs) == 0 {
return nil, nil, fmt.Errorf("%w: task ids are required", ErrTaskCategoryInvalid)
}
if len(taskIDs) > MaxTaskCategoryBatchSize {
return nil, nil, fmt.Errorf("%w: at most %d tasks per request", ErrTaskCategoryInvalid, MaxTaskCategoryBatchSize)
}
for _, id := range taskIDs {
if id <= 0 {
return nil, nil, fmt.Errorf("%w: task id must be positive", ErrTaskCategoryInvalid)
}
}
if categoryID != nil && *categoryID <= 0 {
return nil, nil, fmt.Errorf("%w: category id must be positive", ErrTaskCategoryInvalid)
}
tx, err := d.Begin()
if err != nil {
return nil, nil, err
}
defer tx.Rollback() //nolint:errcheck
// Checking the category inside the transaction keeps a concurrent delete from
// turning the UPDATE below into a foreign key violation.
if categoryID != nil {
var exists bool
if err := tx.QueryRow(`SELECT EXISTS(SELECT 1 FROM task_categories WHERE id=$1)`, *categoryID).Scan(&exists); err != nil {
return nil, nil, err
}
if !exists {
return nil, nil, ErrTaskCategoryNotFound
}
}
rows, err := tx.Query(`
UPDATE tasks SET category_id=$2
WHERE id=ANY($1::bigint[]) AND deleted_at IS NULL
RETURNING id`, taskIDs, categoryID)
if err != nil {
return nil, nil, err
}
updated := make([]int64, 0, len(taskIDs))
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
rows.Close()
return nil, nil, err
}
updated = append(updated, id)
}
rows.Close()
if err := rows.Err(); err != nil {
return nil, nil, err
}
var category *TaskCategory
if categoryID != nil {
fetched, err := scanTaskCategory(tx.QueryRow(`
SELECT `+taskCategoryCols+`
FROM task_categories category
LEFT JOIN tasks task ON task.category_id=category.id
WHERE category.id=$1
GROUP BY category.id`, *categoryID))
if err != nil {
return nil, nil, err
}
category = &fetched
}
if err := tx.Commit(); err != nil {
return nil, nil, err
}
return updated, category, nil
}
+183
View File
@@ -0,0 +1,183 @@
package db
import (
"errors"
"fmt"
"strings"
"testing"
"time"
)
func TestTaskCategoryCRUDAndTaskAssignment(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
name := fmt.Sprintf("Category %d", time.Now().UnixNano())
category, err := d.CreateTaskCategory(" " + strings.ReplaceAll(name, " ", " ") + " ")
if err != nil {
t.Fatal(err)
}
if category.Name != name {
t.Fatalf("category name=%q, want %q", category.Name, name)
}
if _, err := d.CreateTaskCategory(strings.ToUpper(name)); !errors.Is(err, ErrTaskCategoryNameConflict) {
t.Fatalf("duplicate error=%v, want %v", err, ErrTaskCategoryNameConflict)
}
task, err := d.CreateTaskWithOptions("categorized task", "goal", TaskCreateOptions{CategoryID: &category.ID})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(task.ID) })
loaded, err := d.GetTask(task.ID)
if err != nil || loaded == nil || loaded.CategoryID == nil || *loaded.CategoryID != category.ID || loaded.CategoryName != name {
t.Fatalf("categorized task=%+v err=%v", loaded, err)
}
renamed, err := d.RenameTaskCategory(category.ID, name+" renamed")
if err != nil {
t.Fatal(err)
}
if renamed.TaskCount != 1 {
t.Fatalf("renamed task count=%d, want 1", renamed.TaskCount)
}
loaded, _ = d.GetTask(task.ID)
if loaded.CategoryName != renamed.Name {
t.Fatalf("task category name=%q, want %q", loaded.CategoryName, renamed.Name)
}
if _, err := d.SetTaskCategory(task.ID, nil); err != nil {
t.Fatal(err)
}
loaded, _ = d.GetTask(task.ID)
if loaded.CategoryID != nil || loaded.CategoryName != "" {
t.Fatalf("uncategorized task retained category: %+v", loaded)
}
if _, err := d.SetTaskCategory(task.ID, &category.ID); err != nil {
t.Fatal(err)
}
deleted, err := d.DeleteTaskCategory(category.ID)
if err != nil || !deleted {
t.Fatalf("delete category=%v err=%v", deleted, err)
}
loaded, _ = d.GetTask(task.ID)
if loaded.CategoryID != nil || loaded.CategoryName != "" {
t.Fatalf("category deletion did not clear task: %+v", loaded)
}
}
func TestCreateTaskRejectsMissingCategoryAtomically(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
missing := int64(1 << 62)
description := fmt.Sprintf("missing-category-%d", time.Now().UnixNano())
if _, err := d.CreateTaskWithOptions(description, "goal", TaskCreateOptions{CategoryID: &missing}); !errors.Is(err, ErrTaskCategoryNotFound) {
t.Fatalf("create error=%v, want %v", err, ErrTaskCategoryNotFound)
}
var count int
if err := d.QueryRow(`SELECT count(*) FROM explorations WHERE description=$1`, description).Scan(&count); err != nil {
t.Fatal(err)
}
if count != 0 {
t.Fatalf("failed create leaked %d exploration rows", count)
}
}
func TestSetTasksCategoryBatch(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
stamp := time.Now().UnixNano()
category, err := d.CreateTaskCategory(fmt.Sprintf("Batch %d", stamp))
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _, _ = d.DeleteTaskCategory(category.ID) })
ids := make([]int64, 0, 3)
for i := range 3 {
task, err := d.CreateTaskWithOptions(fmt.Sprintf("batch-move-%d-%d", stamp, i), "goal", TaskCreateOptions{})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(task.ID) })
ids = append(ids, task.ID)
}
// A deleted id must not abort the move for the surviving tasks; it is simply
// absent from the returned slice so the caller can report it.
missing := int64(1 << 62)
updated, moved, err := d.SetTasksCategory(append(append([]int64{}, ids...), missing), &category.ID)
if err != nil {
t.Fatal(err)
}
if len(updated) != len(ids) {
t.Fatalf("updated=%v, want the %d live ids only", updated, len(ids))
}
if moved == nil || moved.TaskCount != len(ids) {
t.Fatalf("category=%+v, want task_count=%d", moved, len(ids))
}
for _, id := range ids {
loaded, err := d.GetTask(id)
if err != nil || loaded.CategoryID == nil || *loaded.CategoryID != category.ID {
t.Fatalf("task %d not moved: %+v err=%v", id, loaded, err)
}
}
// A nil category clears the assignment for the whole batch.
updated, moved, err = d.SetTasksCategory(ids, nil)
if err != nil || len(updated) != len(ids) || moved != nil {
t.Fatalf("clear updated=%v category=%+v err=%v", updated, moved, err)
}
for _, id := range ids {
loaded, _ := d.GetTask(id)
if loaded.CategoryID != nil || loaded.CategoryName != "" {
t.Fatalf("task %d still categorized: %+v", id, loaded)
}
}
}
func TestSetTasksCategoryRejectsBadInputAtomically(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
stamp := time.Now().UnixNano()
task, err := d.CreateTaskWithOptions(fmt.Sprintf("batch-reject-%d", stamp), "goal", TaskCreateOptions{})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(task.ID) })
missingCategory := int64(1 << 62)
if _, _, err := d.SetTasksCategory([]int64{task.ID}, &missingCategory); !errors.Is(err, ErrTaskCategoryNotFound) {
t.Fatalf("missing category error=%v, want %v", err, ErrTaskCategoryNotFound)
}
if _, _, err := d.SetTasksCategory(nil, nil); !errors.Is(err, ErrTaskCategoryInvalid) {
t.Fatalf("empty ids error=%v, want %v", err, ErrTaskCategoryInvalid)
}
oversized := make([]int64, MaxTaskCategoryBatchSize+1)
for i := range oversized {
oversized[i] = task.ID
}
if _, _, err := d.SetTasksCategory(oversized, nil); !errors.Is(err, ErrTaskCategoryInvalid) {
t.Fatalf("oversized error=%v, want %v", err, ErrTaskCategoryInvalid)
}
// The rejected calls must leave the task untouched.
loaded, _ := d.GetTask(task.ID)
if loaded.CategoryID != nil {
t.Fatalf("rejected batch mutated task: %+v", loaded)
}
}
+578
View File
@@ -0,0 +1,578 @@
package db
import (
"database/sql"
"fmt"
"strings"
"time"
"unicode/utf8"
)
const IntentBlockedLLMQuota = "llm_quota_exhausted"
// TaskLLMProfile is one ordered entry in a task's explicit failover chain.
type TaskLLMProfile struct {
ProfileID int64 `json:"profile_id"`
Position int `json:"position"`
Status string `json:"status"`
LastError string `json:"last_error,omitempty"`
ExhaustedAt *time.Time `json:"exhausted_at,omitempty"`
}
// TaskSource identifies one directly related task and its exploration.
type TaskSource struct {
TaskID int64
ExplorationID int64
Description string
Goal string
Status string
}
// TaskLLMTransition reports the shared task-level result of marking one profile
// quota-exhausted. NextProfileID is nil when the explicit chain is exhausted.
type TaskLLMTransition struct {
PreviousProfileID int64
NextProfileID *int64
ChainExhausted bool
Advanced bool
// Stale means the chain was replaced after the failing call selected its
// provider. The caller may retry a pre-stream request against the new chain,
// but must not report or persist a transition for this result.
Stale bool
}
func (d *DB) hydrateTaskContext(t *Task) error {
if t == nil {
return nil
}
// A nil cursor on a non-empty chain is a persisted end-of-chain marker. Do not
// repair it from an earlier ready entry: the user may have manually selected a
// profile in the middle and legitimately exhausted every candidate after it.
legacyProfileID, activeProfileID, revision, chain, err := d.taskLLMContext(t.ID)
if err != nil {
return err
}
sources, err := d.TaskSourceIDs(t.ID)
if err != nil {
return err
}
t.SourceTaskIDs = sources
companies, err := d.TaskCompanyIDs(t.ID)
if err != nil {
return err
}
t.CompanyIDs = companies
applyTaskLLMContext(t, legacyProfileID, activeProfileID, revision, chain)
return nil
}
func applyTaskLLMContext(
t *Task,
legacyProfileID, activeProfileID *int64,
revision int64,
chain []TaskLLMProfile,
) {
t.LLMProfileID = legacyProfileID
t.ActiveLLMProfileID = activeProfileID
t.LLMChainRevision = revision
t.LLMProfileIDs = make([]int64, 0, len(chain))
t.LLMFailoverState = "default"
t.LLMFailoverReason = ""
var latest *TaskLLMProfile
activeReady := false
for i := range chain {
entry := chain[i]
t.LLMProfileIDs = append(t.LLMProfileIDs, entry.ProfileID)
if t.ActiveLLMProfileID != nil && entry.ProfileID == *t.ActiveLLMProfileID && entry.Status == "ready" {
activeReady = true
}
if entry.ExhaustedAt != nil && (latest == nil || entry.ExhaustedAt.After(*latest.ExhaustedAt)) {
copy := entry
latest = &copy
}
}
if len(chain) > 0 {
if activeReady {
t.LLMFailoverState = "ready"
} else {
t.LLMFailoverState = "chain_exhausted"
}
}
if latest != nil {
t.LLMFailoverReason = latest.LastError
}
}
type taskBatchContext struct {
legacyProfileID *int64
activeProfileID *int64
revision int64
chain []TaskLLMProfile
sourceIDs []int64
companyIDs []int64
}
// hydrateTasksContext loads every task's LLM chain, source tasks, and company
// scopes with three bulk queries. ListTasks used to issue these queries once per
// task, making startup and task-list hydration grow as 3N+1 database round trips.
func (d *DB) hydrateTasksContext(tasks []*Task) error {
if len(tasks) == 0 {
return nil
}
ids := make([]int64, 0, len(tasks))
contexts := make(map[int64]*taskBatchContext, len(tasks))
byID := make(map[int64]*Task, len(tasks))
for _, task := range tasks {
if task == nil {
continue
}
ids = append(ids, task.ID)
contexts[task.ID] = &taskBatchContext{}
byID[task.ID] = task
}
if len(ids) == 0 {
return nil
}
rows, err := d.Query(`
SELECT t.id, t.llm_profile_id, t.active_llm_profile_id, t.llm_chain_revision,
p.profile_id, p.position, p.status, COALESCE(p.last_error,''), p.exhausted_at
FROM tasks t
LEFT JOIN task_llm_profiles p ON p.task_id=t.id
WHERE t.id=ANY($1::bigint[])
ORDER BY t.id, p.position NULLS LAST`, ids)
if err != nil {
return err
}
for rows.Next() {
var taskID, revision int64
var legacy, active, profileID, position sql.NullInt64
var status, lastError sql.NullString
var exhaustedAt sql.NullTime
if err := rows.Scan(
&taskID, &legacy, &active, &revision,
&profileID, &position, &status, &lastError, &exhaustedAt,
); err != nil {
rows.Close()
return err
}
context := contexts[taskID]
if context == nil {
continue
}
context.revision = revision
if legacy.Valid {
id := legacy.Int64
context.legacyProfileID = &id
}
if active.Valid {
id := active.Int64
context.activeProfileID = &id
}
if profileID.Valid {
entry := TaskLLMProfile{
ProfileID: profileID.Int64,
Position: int(position.Int64),
Status: status.String,
LastError: lastError.String,
}
if exhaustedAt.Valid {
ts := exhaustedAt.Time
entry.ExhaustedAt = &ts
}
context.chain = append(context.chain, entry)
}
}
if err := rows.Err(); err != nil {
rows.Close()
return err
}
rows.Close()
rows, err = d.Query(`
SELECT task_id, source_task_id
FROM task_relations
WHERE task_id=ANY($1::bigint[])
ORDER BY task_id, created_at, source_task_id`, ids)
if err != nil {
return err
}
for rows.Next() {
var taskID, sourceID int64
if err := rows.Scan(&taskID, &sourceID); err != nil {
rows.Close()
return err
}
if context := contexts[taskID]; context != nil {
context.sourceIDs = append(context.sourceIDs, sourceID)
}
}
if err := rows.Err(); err != nil {
rows.Close()
return err
}
rows.Close()
rows, err = d.Query(`
SELECT task_id, company_id
FROM task_scope
WHERE task_id=ANY($1::bigint[]) AND kind='company' AND company_id IS NOT NULL
ORDER BY task_id, id`, ids)
if err != nil {
return err
}
for rows.Next() {
var taskID, companyID int64
if err := rows.Scan(&taskID, &companyID); err != nil {
rows.Close()
return err
}
if context := contexts[taskID]; context != nil {
context.companyIDs = append(context.companyIDs, companyID)
}
}
if err := rows.Err(); err != nil {
rows.Close()
return err
}
rows.Close()
for id, context := range contexts {
task := byID[id]
task.SourceTaskIDs = context.sourceIDs
task.CompanyIDs = context.companyIDs
applyTaskLLMContext(
task,
context.legacyProfileID,
context.activeProfileID,
context.revision,
context.chain,
)
}
return nil
}
// taskLLMContext reads the compatibility profile, current cursor, and ordered
// chain in one statement. Runtime chain edits commit atomically, and one SQL
// statement gives hydration one matching snapshot instead of a transient mix of
// an old task cursor and a newly replaced chain (or vice versa).
func (d *DB) taskLLMContext(taskID int64) (*int64, *int64, int64, []TaskLLMProfile, error) {
rows, err := d.Query(`
SELECT t.llm_profile_id, t.active_llm_profile_id, t.llm_chain_revision,
p.profile_id, p.position, p.status, COALESCE(p.last_error,''), p.exhausted_at
FROM tasks t
LEFT JOIN task_llm_profiles p ON p.task_id=t.id
WHERE t.id=$1
ORDER BY p.position NULLS LAST`, taskID)
if err != nil {
return nil, nil, 0, nil, err
}
defer rows.Close()
var (
legacyProfileID *int64
activeProfileID *int64
revision int64
chain []TaskLLMProfile
found bool
)
for rows.Next() {
var legacy, active, profileID, position sql.NullInt64
var rowRevision int64
var status, lastError sql.NullString
var exhaustedAt sql.NullTime
if err := rows.Scan(&legacy, &active, &rowRevision, &profileID, &position, &status, &lastError, &exhaustedAt); err != nil {
return nil, nil, 0, nil, err
}
if !found {
found = true
revision = rowRevision
if legacy.Valid {
id := legacy.Int64
legacyProfileID = &id
}
if active.Valid {
id := active.Int64
activeProfileID = &id
}
}
if !profileID.Valid {
continue
}
entry := TaskLLMProfile{
ProfileID: profileID.Int64,
Position: int(position.Int64),
Status: status.String,
LastError: lastError.String,
}
if exhaustedAt.Valid {
ts := exhaustedAt.Time
entry.ExhaustedAt = &ts
}
chain = append(chain, entry)
}
if err := rows.Err(); err != nil {
return nil, nil, 0, nil, err
}
if !found {
return nil, nil, 0, nil, sql.ErrNoRows
}
return legacyProfileID, activeProfileID, revision, chain, nil
}
func (d *DB) TaskSourceIDs(taskID int64) ([]int64, error) {
rows, err := d.Query(`SELECT source_task_id FROM task_relations WHERE task_id=$1 ORDER BY created_at, source_task_id`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
var out []int64
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
return nil, err
}
out = append(out, id)
}
return out, rows.Err()
}
// TaskCompanyIDs returns the companies whose asset scopes are available to the
// task. The task_scope rows remain the single source of truth.
func (d *DB) TaskCompanyIDs(taskID int64) ([]int64, error) {
rows, err := d.Query(`
SELECT company_id FROM task_scope
WHERE task_id=$1 AND kind='company' AND company_id IS NOT NULL
ORDER BY id`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
var out []int64
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
return nil, err
}
out = append(out, id)
}
return out, rows.Err()
}
func (d *DB) TaskSources(taskID int64) ([]TaskSource, error) {
rows, err := d.Query(`
SELECT t.id, t.exploration_id, t.description, t.goal, t.status
FROM task_relations r
JOIN tasks t ON t.id=r.source_task_id AND t.deleted_at IS NULL
WHERE r.task_id=$1
ORDER BY r.created_at, r.source_task_id`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
var out []TaskSource
for rows.Next() {
var source TaskSource
if err := rows.Scan(&source.TaskID, &source.ExplorationID, &source.Description, &source.Goal, &source.Status); err != nil {
return nil, err
}
out = append(out, source)
}
return out, rows.Err()
}
func (d *DB) TaskLLMProfiles(taskID int64) ([]TaskLLMProfile, error) {
rows, err := d.Query(`SELECT profile_id, position, status, COALESCE(last_error,''), exhausted_at
FROM task_llm_profiles WHERE task_id=$1 ORDER BY position`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
var out []TaskLLMProfile
for rows.Next() {
var entry TaskLLMProfile
if err := rows.Scan(&entry.ProfileID, &entry.Position, &entry.Status, &entry.LastError, &entry.ExhaustedAt); err != nil {
return nil, err
}
out = append(out, entry)
}
return out, rows.Err()
}
// ReplaceTaskLLMProfiles atomically replaces and resets the explicit task chain.
// activeProfileID=0 selects the first entry. An empty list restores the existing
// agent-binding/global fallback behavior.
// 终态(done/failed/timeout)任务同样允许改链:任务结束后主 Agent 对话仍会走这条链,
// 链上模型不可用时必须能换,否则已完成任务就再也没法交互了。
func (d *DB) ReplaceTaskLLMProfiles(taskID int64, profileIDs []int64, activeProfileID int64) error {
tx, err := d.Begin()
if err != nil {
return err
}
defer tx.Rollback()
// Keep the task -> profile lock order shared by all task LLM mutations and
// DeleteProfile. Inserts below may take KEY SHARE locks on llm_profiles for
// their foreign keys, so the task row must be locked before any of them.
var lockedID int64
if err := tx.QueryRow(`SELECT id FROM tasks WHERE id=$1 AND deleted_at IS NULL FOR UPDATE`, taskID).Scan(&lockedID); err != nil {
if err == sql.ErrNoRows {
return fmt.Errorf("task %d not found", taskID)
}
return err
}
if _, err := tx.Exec(`DELETE FROM task_llm_profiles WHERE task_id=$1`, taskID); err != nil {
return err
}
if err := insertTaskLLMProfiles(tx, taskID, profileIDs); err != nil {
return err
}
if len(profileIDs) == 0 {
activeProfileID = 0
} else if activeProfileID == 0 {
activeProfileID = profileIDs[0]
}
if activeProfileID > 0 {
found := false
for _, id := range profileIDs {
found = found || id == activeProfileID
}
if !found {
return fmt.Errorf("active LLM profile %d is not in the task chain", activeProfileID)
}
}
var active any
if activeProfileID > 0 {
active = activeProfileID
}
if _, err := tx.Exec(`UPDATE tasks
SET active_llm_profile_id=$2, llm_profile_id=$2, llm_chain_revision=llm_chain_revision+1
WHERE id=$1`, taskID, active); err != nil {
return err
}
return tx.Commit()
}
// MarkTaskLLMProfileQuotaExhausted advances the shared task cursor once. A late
// in-flight error from an older profile records that entry as exhausted but does
// not advance past the profile another call already selected.
func (d *DB) MarkTaskLLMProfileQuotaExhausted(taskID, profileID int64, reason string) (TaskLLMTransition, error) {
return d.markTaskLLMProfileQuotaExhausted(taskID, profileID, 0, false, reason)
}
// MarkTaskLLMProfileQuotaExhaustedAtRevision applies a provider failure only
// when it belongs to the chain snapshot used to start that call.
func (d *DB) MarkTaskLLMProfileQuotaExhaustedAtRevision(taskID, profileID, revision int64, reason string) (TaskLLMTransition, error) {
return d.markTaskLLMProfileQuotaExhausted(taskID, profileID, revision, true, reason)
}
func (d *DB) markTaskLLMProfileQuotaExhausted(taskID, profileID, revision int64, checkRevision bool, reason string) (TaskLLMTransition, error) {
var out TaskLLMTransition
tx, err := d.Begin()
if err != nil {
return out, err
}
defer tx.Rollback()
var (
active sql.NullInt64
currentRevision int64
)
if err := tx.QueryRow(`SELECT active_llm_profile_id, llm_chain_revision FROM tasks WHERE id=$1 FOR UPDATE`, taskID).Scan(&active, &currentRevision); err != nil {
return out, err
}
out.PreviousProfileID = profileID
if checkRevision && revision != currentRevision {
out.Stale = true
return out, tx.Commit()
}
var (
position int
entryStatus string
)
if err := tx.QueryRow(`SELECT position, status FROM task_llm_profiles WHERE task_id=$1 AND profile_id=$2`, taskID, profileID).
Scan(&position, &entryStatus); err != nil {
if err == sql.ErrNoRows {
// The chain was edited while this request was in flight. Its provider
// error remains valid for the caller, but it must not mutate the new chain.
return out, tx.Commit()
}
return out, err
}
// Once the cursor reached the end, late failures must be idempotent. This also
// prevents an older in-flight request from reviving a ready entry before a
// manually selected starting position.
if !active.Valid {
out.ChainExhausted = true
return out, tx.Commit()
}
reason = truncateUTF8(strings.TrimSpace(reason), 1000)
if entryStatus != "quota_exhausted" {
if _, err := tx.Exec(`UPDATE task_llm_profiles
SET status='quota_exhausted', last_error=$3, exhausted_at=now()
WHERE task_id=$1 AND profile_id=$2`, taskID, profileID, reason); err != nil {
return out, err
}
}
if active.Valid && active.Int64 != profileID {
next := active.Int64
out.NextProfileID = &next
return out, tx.Commit()
}
var next int64
err = tx.QueryRow(`SELECT profile_id FROM task_llm_profiles
WHERE task_id=$1 AND position>$2 AND status='ready'
ORDER BY position LIMIT 1`, taskID, position).Scan(&next)
switch err {
case nil:
out.Advanced = true
out.NextProfileID = &next
if _, err := tx.Exec(`UPDATE tasks
SET active_llm_profile_id=$2, llm_profile_id=$2, llm_chain_revision=llm_chain_revision+1
WHERE id=$1`, taskID, next); err != nil {
return out, err
}
case sql.ErrNoRows:
out.Advanced = true
out.ChainExhausted = true
if _, err := tx.Exec(`UPDATE tasks
SET active_llm_profile_id=NULL, llm_profile_id=NULL, llm_chain_revision=llm_chain_revision+1
WHERE id=$1`, taskID); err != nil {
return out, err
}
default:
return out, err
}
return out, tx.Commit()
}
func truncateUTF8(value string, maxBytes int) string {
value = strings.ToValidUTF8(value, "\uFFFD")
if maxBytes <= 0 {
return ""
}
if len(value) <= maxBytes {
return value
}
end := maxBytes
for end > 0 && !utf8.ValidString(value[:end]) {
end--
}
return value[:end]
}
func (s *ExplorationStore) SetIntentBlockedReason(id int64, reason string) error {
_, err := s.db.Exec(`UPDATE exploration_nodes
SET state='blocked', blocked_reason=NULLIF($1,''), completed_at=now()
WHERE id=$2 AND exploration_id=$3 AND kind='intent'`, reason, id, s.expID)
return err
}
func (s *ExplorationStore) ReopenIntentsByBlockedReason(reason string) (int64, error) {
res, err := s.db.Exec(`UPDATE exploration_nodes
SET state='open', blocked_reason=NULL, completed_at=NULL
WHERE exploration_id=$1 AND kind='intent' AND state='blocked' AND blocked_reason=$2`, s.expID, reason)
if err != nil {
return 0, err
}
n, _ := res.RowsAffected()
return n, nil
}
+527
View File
@@ -0,0 +1,527 @@
package db
import (
"database/sql"
"fmt"
"testing"
"time"
)
func TestTaskLLMProfileMutationsLockTaskBeforeProfile(t *testing.T) {
dsn := testDSN(t)
d, err := Open(dsn)
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
suffix := time.Now().UnixNano()
first, err := d.SaveProfile(&LLMProfile{
Name: fmt.Sprintf("lock-order-first-%d", suffix), Format: "openai", Model: "first", APIKey: "test-key",
})
if err != nil {
t.Fatal(err)
}
second, err := d.SaveProfile(&LLMProfile{
Name: fmt.Sprintf("lock-order-second-%d", suffix), Format: "openai", Model: "second", APIKey: "test-key",
})
if err != nil {
_ = d.DeleteProfile(first)
t.Fatal(err)
}
task, err := d.CreateTaskWithOptions("LLM lock order", "verify concurrent mutation locks", TaskCreateOptions{
LLMProfileIDs: []int64{first, second},
})
if err != nil {
_ = d.DeleteProfile(first)
_ = d.DeleteProfile(second)
t.Fatal(err)
}
t.Cleanup(func() {
_ = d.DeleteTask(task.ID)
_ = d.DeleteProfile(first)
_ = d.DeleteProfile(second)
})
replaceDB, replacePID := openSingleConnectionTestDB(t, dsn)
replaceBlocker, err := d.Begin()
if err != nil {
t.Fatal(err)
}
defer replaceBlocker.Rollback()
if _, err := replaceBlocker.Exec(`SELECT id FROM tasks WHERE id=$1 FOR UPDATE`, task.ID); err != nil {
replaceBlocker.Rollback()
t.Fatal(err)
}
replaceDone := make(chan error, 1)
go func() {
replaceDone <- replaceDB.ReplaceTaskLLMProfiles(task.ID, []int64{second, first}, second)
}()
if err := waitForBackendBlock(d, replacePID, replaceDone); err != nil {
replaceBlocker.Rollback()
t.Fatal(err)
}
assertProfilesUnlocked(t, d, first, second)
if err := replaceBlocker.Rollback(); err != nil {
t.Fatal(err)
}
if err := waitForMutationResult(replaceDone); err != nil {
t.Fatalf("replace chain: %v", err)
}
deleteDB, deletePID := openSingleConnectionTestDB(t, dsn)
deleteBlocker, err := d.Begin()
if err != nil {
t.Fatal(err)
}
defer deleteBlocker.Rollback()
if _, err := deleteBlocker.Exec(`SELECT id FROM tasks WHERE id=$1 FOR UPDATE`, task.ID); err != nil {
deleteBlocker.Rollback()
t.Fatal(err)
}
deleteDone := make(chan error, 1)
go func() {
deleteDone <- deleteDB.DeleteProfile(second)
}()
if err := waitForBackendBlock(d, deletePID, deleteDone); err != nil {
deleteBlocker.Rollback()
t.Fatal(err)
}
assertProfilesUnlocked(t, d, second)
if err := deleteBlocker.Rollback(); err != nil {
t.Fatal(err)
}
if err := waitForMutationResult(deleteDone); err != nil {
t.Fatalf("delete profile: %v", err)
}
got, err := d.GetTask(task.ID)
if err != nil {
t.Fatal(err)
}
if got.ActiveLLMProfileID == nil || *got.ActiveLLMProfileID != first {
t.Fatalf("deleting the active profile did not select its successor: %+v", got)
}
}
func TestDeleteProfileLocksNonTaskReferencesBeforeProfile(t *testing.T) {
dsn := testDSN(t)
t.Run("agent", func(t *testing.T) {
d, err := Open(dsn)
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
suffix := time.Now().UnixNano()
profileID, err := d.SaveProfile(&LLMProfile{
Name: fmt.Sprintf("agent-lock-profile-%d", suffix), Format: "openai", Model: "agent-lock", APIKey: "test-key",
})
if err != nil {
t.Fatal(err)
}
agent, err := d.CreateAgent(fmt.Sprintf("lock_agent_%d", suffix), "lock agent", "")
if err != nil {
_ = d.DeleteProfile(profileID)
t.Fatal(err)
}
if err := d.SetAgentLLMProfile(agent.Key, &profileID); err != nil {
_ = d.DeleteAgent(agent.Key)
_ = d.DeleteProfile(profileID)
t.Fatal(err)
}
t.Cleanup(func() {
_ = d.DeleteAgent(agent.Key)
_ = d.DeleteProfile(profileID)
})
blocker, err := d.Begin()
if err != nil {
t.Fatal(err)
}
defer blocker.Rollback()
if _, err := blocker.Exec(`SELECT id FROM agents WHERE id=$1 FOR UPDATE`, agent.ID); err != nil {
t.Fatal(err)
}
deleteDB, deletePID := openSingleConnectionTestDB(t, dsn)
deleteDone := make(chan error, 1)
go func() { deleteDone <- deleteDB.DeleteProfile(profileID) }()
if err := waitForBackendBlock(d, deletePID, deleteDone); err != nil {
blocker.Rollback()
t.Fatal(err)
}
assertProfilesUnlocked(t, d, profileID)
if err := blocker.Rollback(); err != nil {
t.Fatal(err)
}
if err := waitForMutationResult(deleteDone); err != nil {
t.Fatalf("delete agent profile: %v", err)
}
got, err := d.GetAgentByKey(agent.Key)
if err != nil || got == nil || got.LLMProfileID != nil {
t.Fatalf("agent binding was not cleared: agent=%+v err=%v", got, err)
}
})
t.Run("conversation", func(t *testing.T) {
d, err := Open(dsn)
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
suffix := time.Now().UnixNano()
profileID, err := d.SaveProfile(&LLMProfile{
Name: fmt.Sprintf("conversation-lock-profile-%d", suffix), Format: "openai", Model: "conversation-lock", APIKey: "test-key",
})
if err != nil {
t.Fatal(err)
}
conversation, err := d.CreateConversation("planner", "lock conversation", &profileID)
if err != nil {
_ = d.DeleteProfile(profileID)
t.Fatal(err)
}
t.Cleanup(func() {
_ = d.DeleteConversation(conversation.ID)
_ = d.DeleteProfile(profileID)
})
blocker, err := d.Begin()
if err != nil {
t.Fatal(err)
}
defer blocker.Rollback()
if _, err := blocker.Exec(`SELECT id FROM conversations WHERE id=$1 FOR UPDATE`, conversation.ID); err != nil {
t.Fatal(err)
}
deleteDB, deletePID := openSingleConnectionTestDB(t, dsn)
deleteDone := make(chan error, 1)
go func() { deleteDone <- deleteDB.DeleteProfile(profileID) }()
if err := waitForBackendBlock(d, deletePID, deleteDone); err != nil {
blocker.Rollback()
t.Fatal(err)
}
assertProfilesUnlocked(t, d, profileID)
if err := blocker.Rollback(); err != nil {
t.Fatal(err)
}
if err := waitForMutationResult(deleteDone); err != nil {
t.Fatalf("delete conversation profile: %v", err)
}
got, err := d.GetConversation(conversation.ID)
if err != nil || got == nil || got.LLMProfileID != nil {
t.Fatalf("conversation binding was not cleared: conversation=%+v err=%v", got, err)
}
})
}
func TestNonTaskProfileMutationsLockReferenceBeforeProfile(t *testing.T) {
dsn := testDSN(t)
t.Run("agent", func(t *testing.T) {
d, err := Open(dsn)
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
suffix := time.Now().UnixNano()
profileID, err := d.SaveProfile(&LLMProfile{
Name: fmt.Sprintf("agent-write-lock-profile-%d", suffix), Format: "openai", Model: "agent-write-lock", APIKey: "test-key",
})
if err != nil {
t.Fatal(err)
}
agent, err := d.CreateAgent(fmt.Sprintf("write_lock_agent_%d", suffix), "write lock agent", "")
if err != nil {
_ = d.DeleteProfile(profileID)
t.Fatal(err)
}
t.Cleanup(func() {
_ = d.DeleteAgent(agent.Key)
_ = d.DeleteProfile(profileID)
})
blocker, err := d.Begin()
if err != nil {
t.Fatal(err)
}
defer blocker.Rollback()
if _, err := blocker.Exec(`SELECT id FROM agents WHERE id=$1 FOR UPDATE`, agent.ID); err != nil {
t.Fatal(err)
}
mutationDB, mutationPID := openSingleConnectionTestDB(t, dsn)
mutationDone := make(chan error, 1)
go func() { mutationDone <- mutationDB.SetAgentLLMProfile(agent.Key, &profileID) }()
if err := waitForBackendBlock(d, mutationPID, mutationDone); err != nil {
blocker.Rollback()
t.Fatal(err)
}
assertProfilesUnlocked(t, d, profileID)
if err := blocker.Rollback(); err != nil {
t.Fatal(err)
}
if err := waitForMutationResult(mutationDone); err != nil {
t.Fatalf("bind agent profile: %v", err)
}
})
t.Run("conversation", func(t *testing.T) {
d, err := Open(dsn)
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
suffix := time.Now().UnixNano()
profileID, err := d.SaveProfile(&LLMProfile{
Name: fmt.Sprintf("conversation-write-lock-profile-%d", suffix), Format: "openai", Model: "conversation-write-lock", APIKey: "test-key",
})
if err != nil {
t.Fatal(err)
}
conversation, err := d.CreateConversation("planner", "write lock conversation", nil)
if err != nil {
_ = d.DeleteProfile(profileID)
t.Fatal(err)
}
t.Cleanup(func() {
_ = d.DeleteConversation(conversation.ID)
_ = d.DeleteProfile(profileID)
})
blocker, err := d.Begin()
if err != nil {
t.Fatal(err)
}
defer blocker.Rollback()
if _, err := blocker.Exec(`SELECT id FROM conversations WHERE id=$1 FOR UPDATE`, conversation.ID); err != nil {
t.Fatal(err)
}
mutationDB, mutationPID := openSingleConnectionTestDB(t, dsn)
mutationDone := make(chan error, 1)
go func() { mutationDone <- mutationDB.UpdateConversationProfile(conversation.ID, &profileID) }()
if err := waitForBackendBlock(d, mutationPID, mutationDone); err != nil {
blocker.Rollback()
t.Fatal(err)
}
assertProfilesUnlocked(t, d, profileID)
if err := blocker.Rollback(); err != nil {
t.Fatal(err)
}
if err := waitForMutationResult(mutationDone); err != nil {
t.Fatalf("bind conversation profile: %v", err)
}
})
}
func TestCreateConversationAndDeleteProfileDoNotDeadlock(t *testing.T) {
dsn := testDSN(t)
d, err := Open(dsn)
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
suffix := time.Now().UnixNano()
profileID, err := d.SaveProfile(&LLMProfile{
Name: fmt.Sprintf("conversation-create-race-profile-%d", suffix), Format: "openai", Model: "conversation-create-race", APIKey: "test-key",
})
if err != nil {
t.Fatal(err)
}
title := fmt.Sprintf("conversation create race %d", suffix)
t.Cleanup(func() {
_, _ = d.Exec(`DELETE FROM conversations WHERE title=$1`, title)
_ = d.DeleteProfile(profileID)
})
profileBlocker, err := d.Begin()
if err != nil {
t.Fatal(err)
}
defer profileBlocker.Rollback()
if _, err := profileBlocker.Exec(`SELECT id FROM llm_profiles WHERE id=$1 FOR UPDATE`, profileID); err != nil {
t.Fatal(err)
}
createDB, createPID := openSingleConnectionTestDB(t, dsn)
type createResult struct {
conversation *Conversation
err error
}
createDone := make(chan createResult, 1)
createStatus := make(chan error, 1)
go func() {
conversation, err := createDB.CreateConversation("planner", title, &profileID)
createDone <- createResult{conversation: conversation, err: err}
createStatus <- err
}()
if err := waitForBackendBlock(d, createPID, createStatus); err != nil {
profileBlocker.Rollback()
t.Fatal(err)
}
deleteDB, deletePID := openSingleConnectionTestDB(t, dsn)
deleteDone := make(chan error, 1)
go func() { deleteDone <- deleteDB.DeleteProfile(profileID) }()
if err := waitForBackendBlock(d, deletePID, deleteDone); err != nil {
profileBlocker.Rollback()
t.Fatal(err)
}
if err := profileBlocker.Rollback(); err != nil {
t.Fatal(err)
}
var created createResult
select {
case created = <-createDone:
case <-time.After(12 * time.Second):
t.Fatal("timed out waiting for conversation creation")
}
if err := waitForMutationResult(deleteDone); err != nil {
t.Fatalf("delete profile during conversation creation: %v", err)
}
var persisted int
if err := d.QueryRow(`SELECT count(*) FROM conversations WHERE title=$1`, title).Scan(&persisted); err != nil {
t.Fatal(err)
}
if created.err == nil {
if created.conversation == nil || persisted != 1 {
t.Fatalf("successful creation was not committed atomically: conversation=%+v count=%d", created.conversation, persisted)
}
} else if persisted != 0 {
t.Fatalf("failed creation left a partial conversation row: err=%v count=%d", created.err, persisted)
}
}
func TestDeleteProfileRetriesReferenceCommittedAfterInitialScan(t *testing.T) {
dsn := testDSN(t)
d, err := Open(dsn)
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
suffix := time.Now().UnixNano()
profileID, err := d.SaveProfile(&LLMProfile{
Name: fmt.Sprintf("late-reference-profile-%d", suffix), Format: "openai", Model: "late-reference", APIKey: "test-key",
})
if err != nil {
t.Fatal(err)
}
task, err := d.CreateTask("late profile reference", "exercise delete retry", nil, 0, 0)
if err != nil {
_ = d.DeleteProfile(profileID)
t.Fatal(err)
}
t.Cleanup(func() {
_ = d.DeleteTask(task.ID)
_ = d.DeleteProfile(profileID)
})
// Keep the new reference uncommitted while deletion takes its initial
// READ COMMITTED snapshot. The task is therefore absent from the first lock
// set, while its FK KEY SHARE lock makes deletion wait at the profile row.
referenceTx, err := d.Begin()
if err != nil {
t.Fatal(err)
}
defer referenceTx.Rollback()
if _, err := referenceTx.Exec(`UPDATE tasks
SET llm_profile_id=$2, active_llm_profile_id=$2
WHERE id=$1`, task.ID, profileID); err != nil {
t.Fatal(err)
}
if _, err := referenceTx.Exec(`INSERT INTO task_llm_profiles(task_id, profile_id, position)
VALUES ($1,$2,0)`, task.ID, profileID); err != nil {
t.Fatal(err)
}
deleteDB, deletePID := openSingleConnectionTestDB(t, dsn)
deleteDone := make(chan error, 1)
go func() { deleteDone <- deleteDB.DeleteProfile(profileID) }()
if err := waitForBackendBlock(d, deletePID, deleteDone); err != nil {
referenceTx.Rollback()
t.Fatal(err)
}
if err := referenceTx.Commit(); err != nil {
t.Fatal(err)
}
if err := waitForMutationResult(deleteDone); err != nil {
t.Fatalf("delete profile after late reference: %v", err)
}
got, err := d.GetTask(task.ID)
if err != nil {
t.Fatal(err)
}
if got.ActiveLLMProfileID != nil || len(got.LLMProfileIDs) != 0 || got.LLMChainRevision != 1 {
t.Fatalf("late task reference was not handled by a locked retry: %+v", got)
}
}
func openSingleConnectionTestDB(t *testing.T, dsn string) (*DB, int) {
t.Helper()
sqlDB, err := sql.Open("pgx", dsn)
if err != nil {
t.Fatal(err)
}
sqlDB.SetMaxOpenConns(1)
sqlDB.SetMaxIdleConns(1)
t.Cleanup(func() { _ = sqlDB.Close() })
if _, err := sqlDB.Exec(`SET statement_timeout='10s'`); err != nil {
t.Fatal(err)
}
var pid int
if err := sqlDB.QueryRow(`SELECT pg_backend_pid()`).Scan(&pid); err != nil {
t.Fatal(err)
}
return &DB{sqlDB}, pid
}
func waitForBackendBlock(observer *DB, pid int, done <-chan error) error {
deadline := time.Now().Add(5 * time.Second)
for time.Now().Before(deadline) {
select {
case err := <-done:
return fmt.Errorf("mutation returned before reaching the expected reference-row lock: %v", err)
default:
}
var blockers int
if err := observer.QueryRow(`SELECT cardinality(pg_blocking_pids($1))`, pid).Scan(&blockers); err != nil {
return err
}
if blockers > 0 {
return nil
}
time.Sleep(10 * time.Millisecond)
}
return fmt.Errorf("backend %d did not block within 5s", pid)
}
func assertProfilesUnlocked(t *testing.T, d *DB, profileIDs ...int64) {
t.Helper()
probe, err := d.Begin()
if err != nil {
t.Fatal(err)
}
defer probe.Rollback()
for _, profileID := range profileIDs {
var lockedID int64
if err := probe.QueryRow(`SELECT id FROM llm_profiles WHERE id=$1 FOR UPDATE NOWAIT`, profileID).Scan(&lockedID); err != nil {
t.Fatalf("profile %d was locked before the task row: %v", profileID, err)
}
}
}
func waitForMutationResult(done <-chan error) error {
select {
case err := <-done:
return err
case <-time.After(12 * time.Second):
return fmt.Errorf("timed out waiting for task LLM mutation")
}
}
+27
View File
@@ -0,0 +1,27 @@
package db
import (
"strings"
"testing"
"unicode/utf8"
)
func TestTruncateUTF8PreservesValidEncoding(t *testing.T) {
t.Parallel()
input := strings.Repeat("额度不足", 400)
got := truncateUTF8(input, 1000)
if len(got) > 1000 {
t.Fatalf("truncated value has %d bytes, want at most 1000", len(got))
}
if !utf8.ValidString(got) {
t.Fatalf("truncated value is not valid UTF-8: %q", got[len(got)-8:])
}
}
func TestTruncateUTF8RepairsInvalidInput(t *testing.T) {
t.Parallel()
got := truncateUTF8("bad\xffvalue", 1000)
if !utf8.ValidString(got) {
t.Fatalf("repaired value is not valid UTF-8: %q", got)
}
}
+210
View File
@@ -0,0 +1,210 @@
package db
import (
"database/sql"
"fmt"
"testing"
"time"
_ "github.com/jackc/pgx/v5/stdlib"
)
func TestDeleteTaskTrafficHostsUseOneLockedTransaction(t *testing.T) {
dsn := testDSN(t)
d, err := Open(dsn)
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
t.Run("sharing committed before deletion lock is observed", func(t *testing.T) {
first, second, host, rootAssetID := createTaskDeleteRaceFixture(t, d)
serviceURL := "https://" + host + "/concurrent-owner"
t.Cleanup(func() {
_ = d.DeleteTask(first.ID)
_ = d.DeleteTask(second.ID)
_, _ = d.Exec(`DELETE FROM assets WHERE id=$1 OR url=$2`, rootAssetID, serviceURL)
})
writer, _ := openTaskDeleteTestDB(t, dsn)
writerTx, err := writer.Begin()
if err != nil {
t.Fatal(err)
}
defer writerTx.Rollback()
if _, err := writerTx.Exec(`
INSERT INTO assets(type, url, service_type, domain, task_ids)
VALUES ('service', $1, 'http', $2, ARRAY[$3]::bigint[])`, serviceURL, host, second.ID); err != nil {
t.Fatal(err)
}
deleter, deleterPID := openTaskDeleteTestDB(t, dsn)
var preparedHosts []string
deleteDone := make(chan error, 1)
go func() {
_, deleteErr := deleter.DeleteTaskCascadePrepared(first.ID, true, false, false, func(p TaskDeletePreparation) error {
preparedHosts = append([]string(nil), p.TrafficHosts...)
return nil
})
deleteDone <- deleteErr
}()
if err := waitForTaskDeleteBlock(d, deleterPID, deleteDone); err != nil {
t.Fatal(err)
}
if err := writerTx.Commit(); err != nil {
t.Fatal(err)
}
if err := waitForTaskDeleteResult(deleteDone); err != nil {
t.Fatal(err)
}
if containsDeleteHost(preparedHosts, host) {
t.Fatalf("newly shared host %q was selected for traffic deletion: %v", host, preparedHosts)
}
var remaining int
if err := d.QueryRow(`SELECT count(*) FROM assets WHERE url=$1 AND $2=ANY(task_ids)`, serviceURL, second.ID).Scan(&remaining); err != nil {
t.Fatal(err)
}
if remaining != 1 {
t.Fatalf("concurrent owner's asset was not preserved: count=%d", remaining)
}
})
t.Run("sharing cannot enter between host resolution and commit", func(t *testing.T) {
first, second, host, rootAssetID := createTaskDeleteRaceFixture(t, d)
serviceURL := "https://" + host + "/late-owner"
t.Cleanup(func() {
_ = d.DeleteTask(first.ID)
_ = d.DeleteTask(second.ID)
_, _ = d.Exec(`DELETE FROM assets WHERE id=$1 OR url=$2`, rootAssetID, serviceURL)
})
deleter, _ := openTaskDeleteTestDB(t, dsn)
prepared := make(chan TaskDeletePreparation, 1)
releasePrepare := make(chan struct{})
deleteDone := make(chan error, 1)
go func() {
_, deleteErr := deleter.DeleteTaskCascadePrepared(first.ID, true, false, false, func(p TaskDeletePreparation) error {
prepared <- p
<-releasePrepare
return nil
})
deleteDone <- deleteErr
}()
var plan TaskDeletePreparation
select {
case plan = <-prepared:
case err := <-deleteDone:
close(releasePrepare)
t.Fatalf("delete returned before preparation: %v", err)
case <-time.After(5 * time.Second):
close(releasePrepare)
t.Fatal("timed out waiting for task delete preparation")
}
if !containsDeleteHost(plan.TrafficHosts, host) {
close(releasePrepare)
t.Fatalf("exclusive host %q missing from preparation: %v", host, plan.TrafficHosts)
}
writer, writerPID := openTaskDeleteTestDB(t, dsn)
writerDone := make(chan error, 1)
go func() {
_, writeErr := writer.Exec(`
INSERT INTO assets(type, url, service_type, domain, task_ids)
VALUES ('service', $1, 'http', $2, ARRAY[$3]::bigint[])`, serviceURL, host, second.ID)
writerDone <- writeErr
}()
if err := waitForTaskDeleteBlock(d, writerPID, writerDone); err != nil {
close(releasePrepare)
t.Fatal(err)
}
close(releasePrepare)
if err := waitForTaskDeleteResult(deleteDone); err != nil {
t.Fatal(err)
}
if err := waitForTaskDeleteResult(writerDone); err != nil {
t.Fatal(err)
}
})
}
func createTaskDeleteRaceFixture(t *testing.T, d *DB) (first, second *Task, host string, rootAssetID int64) {
t.Helper()
var err error
first, err = d.CreateTask("delete race owner", "delete safely", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
second, err = d.CreateTask("delete race sharer", "preserve shared host", nil, 0, 0)
if err != nil {
_ = d.DeleteTask(first.ID)
t.Fatal(err)
}
host = fmt.Sprintf("task-delete-race-%d.example.test", first.ID)
rootAssetID, err = d.Assets().UpsertRootDomain(UpsertRootDomainReq{Domain: host, TaskID: first.ID})
if err != nil {
_ = d.DeleteTask(first.ID)
_ = d.DeleteTask(second.ID)
t.Fatal(err)
}
return first, second, host, rootAssetID
}
func openTaskDeleteTestDB(t *testing.T, dsn string) (*DB, int) {
t.Helper()
sqlDB, err := sql.Open("pgx", dsn)
if err != nil {
t.Fatal(err)
}
sqlDB.SetMaxOpenConns(1)
sqlDB.SetMaxIdleConns(1)
t.Cleanup(func() { _ = sqlDB.Close() })
if _, err := sqlDB.Exec(`SET statement_timeout='10s'`); err != nil {
t.Fatal(err)
}
var pid int
if err := sqlDB.QueryRow(`SELECT pg_backend_pid()`).Scan(&pid); err != nil {
t.Fatal(err)
}
return &DB{sqlDB}, pid
}
func waitForTaskDeleteBlock(observer *DB, pid int, done <-chan error) error {
deadline := time.Now().Add(5 * time.Second)
for time.Now().Before(deadline) {
select {
case err := <-done:
return fmt.Errorf("operation returned before reaching the task deletion lock: %v", err)
default:
}
var blockers int
if err := observer.QueryRow(`SELECT cardinality(pg_blocking_pids($1))`, pid).Scan(&blockers); err != nil {
return err
}
if blockers > 0 {
return nil
}
time.Sleep(10 * time.Millisecond)
}
return fmt.Errorf("backend %d did not block within 5s", pid)
}
func waitForTaskDeleteResult(done <-chan error) error {
select {
case err := <-done:
return err
case <-time.After(12 * time.Second):
return fmt.Errorf("timed out waiting for concurrent task deletion operation")
}
}
func containsDeleteHost(hosts []string, want string) bool {
for _, host := range hosts {
if host == want {
return true
}
}
return false
}
+122
View File
@@ -0,0 +1,122 @@
package db
import (
"database/sql"
"fmt"
)
// TaskInterceptRuleInput is one task-level rule supplied at task creation.
// Action: 'block'=拦截 'allow'=允许(白名单);空视为 'block'。
type TaskInterceptRuleInput struct {
Enabled bool `json:"enabled"`
Action string `json:"action"`
Kind string `json:"kind"`
Pattern string `json:"pattern"`
Note string `json:"note"`
}
const taskInterceptRuleCols = `id, enabled, action, kind, pattern, note, created_at, updated_at`
func scanTaskInterceptRule(row interface{ Scan(...any) error }) (AssetInterceptRule, error) {
var r AssetInterceptRule
err := row.Scan(&r.ID, &r.Enabled, &r.Action, &r.Kind, &r.Pattern, &r.Note, &r.CreatedAt, &r.UpdatedAt)
return r, err
}
// ListTaskInterceptRules returns a task's rules (both block and allow) as
// AssetInterceptRule (Builtin always false; Action carries block/allow). taskID
// <= 0 returns nothing.
func (s *AssetStore) ListTaskInterceptRules(taskID int64) ([]AssetInterceptRule, error) {
if taskID <= 0 {
return nil, nil
}
rows, err := s.db.Query(`SELECT `+taskInterceptRuleCols+` FROM task_intercept_rules WHERE task_id=$1 ORDER BY action, id`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
var out []AssetInterceptRule
for rows.Next() {
r, err := scanTaskInterceptRule(rows)
if err != nil {
return nil, err
}
out = append(out, r)
}
return out, rows.Err()
}
// TaskInterceptRulesSplit loads a task's rules and splits them into block and
// allow sets, for the enforcement gate.
func (s *AssetStore) TaskInterceptRulesSplit(taskID int64) (block, allow []AssetInterceptRule, err error) {
rules, err := s.ListTaskInterceptRules(taskID)
if err != nil {
return nil, nil, err
}
for _, r := range rules {
if r.Action == "allow" {
allow = append(allow, r)
} else {
block = append(block, r)
}
}
return block, allow, nil
}
func normalizeRuleAction(action string) string {
if action == "allow" {
return "allow"
}
return "block"
}
// CreateTaskInterceptRule inserts a rule under a task.
func (s *AssetStore) CreateTaskInterceptRule(taskID int64, action, kind, pattern, note string, enabled bool) (AssetInterceptRule, error) {
row := s.db.QueryRow(`
INSERT INTO task_intercept_rules(task_id, enabled, action, kind, pattern, note)
VALUES ($1,$2,$3,$4,$5,$6)
RETURNING `+taskInterceptRuleCols,
taskID, enabled, normalizeRuleAction(action), kind, pattern, note)
return scanTaskInterceptRule(row)
}
// UpdateTaskInterceptRule replaces the editable fields of a task's rule (scoped
// by task_id so a rule can only be edited through its owning task).
func (s *AssetStore) UpdateTaskInterceptRule(taskID, ruleID int64, action, kind, pattern, note string, enabled bool) (AssetInterceptRule, error) {
row := s.db.QueryRow(`
UPDATE task_intercept_rules
SET enabled=$3, action=$4, kind=$5, pattern=$6, note=$7
WHERE id=$1 AND task_id=$2
RETURNING `+taskInterceptRuleCols,
ruleID, taskID, enabled, normalizeRuleAction(action), kind, pattern, note)
return scanTaskInterceptRule(row)
}
// DeleteTaskInterceptRule removes a task's rule. Returns false if not found.
func (s *AssetStore) DeleteTaskInterceptRule(taskID, ruleID int64) (bool, error) {
res, err := s.db.Exec(`DELETE FROM task_intercept_rules WHERE id=$1 AND task_id=$2`, ruleID, taskID)
if err != nil {
return false, err
}
n, _ := res.RowsAffected()
return n > 0, nil
}
// ToggleTaskInterceptRule flips the enabled state of a task's rule.
func (s *AssetStore) ToggleTaskInterceptRule(taskID, ruleID int64, enabled bool) error {
_, err := s.db.Exec(`UPDATE task_intercept_rules SET enabled=$3 WHERE id=$1 AND task_id=$2`, ruleID, taskID, enabled)
return err
}
// insertTaskInterceptRules inserts task-level rules within the task-creation
// transaction (mirrors insertTaskCompanies).
func insertTaskInterceptRules(tx *sql.Tx, taskID int64, rules []TaskInterceptRuleInput) error {
for _, r := range rules {
if _, err := tx.Exec(`
INSERT INTO task_intercept_rules(task_id, enabled, action, kind, pattern, note)
VALUES ($1,$2,$3,$4,$5,$6)`, taskID, r.Enabled, normalizeRuleAction(r.Action), r.Kind, r.Pattern, r.Note); err != nil {
return fmt.Errorf("insert task intercept rule %q: %w", r.Pattern, err)
}
}
return nil
}
+117
View File
@@ -0,0 +1,117 @@
package db
import (
"fmt"
"testing"
"time"
)
func TestListTasksBulkHydratesTaskContext(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
stamp := time.Now().UnixNano()
profileID, err := d.SaveProfile(&LLMProfile{
Name: fmt.Sprintf("bulk-list-profile-%d", stamp), Format: "openai", Model: "test", APIKey: "test-key",
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteProfile(profileID) })
source, err := d.CreateTask(fmt.Sprintf("bulk-list-source-%d", stamp), "goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(source.ID) })
companyID, _, err := d.Companies().UpsertCompany(fmt.Sprintf("Bulk List Company %d", stamp), "")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.Companies().DeleteCompany(companyID) })
child, err := d.CreateTaskWithOptions("bulk-list-child", "goal", TaskCreateOptions{
LLMProfileIDs: []int64{profileID},
SourceTaskIDs: []int64{source.ID},
CompanyIDs: []int64{companyID},
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(child.ID) })
tasks, err := d.ListTasks()
if err != nil {
t.Fatal(err)
}
var got *Task
for _, task := range tasks {
if task.ID == child.ID {
got = task
break
}
}
if got == nil {
t.Fatalf("task %d missing from list", child.ID)
}
if len(got.LLMProfileIDs) != 1 || got.LLMProfileIDs[0] != profileID ||
got.ActiveLLMProfileID == nil || *got.ActiveLLMProfileID != profileID || got.LLMFailoverState != "ready" {
t.Fatalf("LLM context was not bulk hydrated: %+v", got)
}
if len(got.SourceTaskIDs) != 1 || got.SourceTaskIDs[0] != source.ID {
t.Fatalf("source task context was not bulk hydrated: %v", got.SourceTaskIDs)
}
if len(got.CompanyIDs) != 1 || got.CompanyIDs[0] != companyID {
t.Fatalf("company context was not bulk hydrated: %v", got.CompanyIDs)
}
}
func TestTaskListMetricsAll(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
task, err := d.CreateTask("task-list-metrics", "goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(task.ID) })
resultAt := time.Now().Add(-time.Second).Truncate(time.Second)
latestAt := resultAt.Add(time.Second)
if _, err := d.Exec(`
INSERT INTO activity(
exploration_id, kind, input_tokens, output_tokens,
cache_read_tokens, cache_write_tokens, created_at
) VALUES ($1,'result',11,7,3,2,$2), ($1,'tool_result',999,999,999,999,$3)`,
task.ExplorationID, resultAt, latestAt); err != nil {
t.Fatal(err)
}
if _, err := d.Exec(`
INSERT INTO exploration_nodes(exploration_id, kind, payload, state)
VALUES ($1,'goal','{}','met'), ($1,'goal','{}','open')`, task.ExplorationID); err != nil {
t.Fatal(err)
}
all, err := d.TaskListMetricsAll()
if err != nil {
t.Fatal(err)
}
metrics, ok := all[task.ExplorationID]
if !ok {
t.Fatalf("metrics for exploration %d missing", task.ExplorationID)
}
if metrics.Tokens.InputTokens != 11 || metrics.Tokens.OutputTokens != 7 ||
metrics.Tokens.CacheReadTokens != 3 || metrics.Tokens.CacheWriteTokens != 2 {
t.Fatalf("unexpected token metrics: %+v", metrics.Tokens)
}
if metrics.LastActivity != latestAt.Unix() {
t.Fatalf("last activity=%d, want %d", metrics.LastActivity, latestAt.Unix())
}
if metrics.Goals.Total != 2 || metrics.Goals.Met != 1 {
t.Fatalf("unexpected goal metrics: %+v", metrics.Goals)
}
}
+104
View File
@@ -0,0 +1,104 @@
package db
import (
"fmt"
"testing"
"time"
)
func TestTaskPinOrderingAndRename(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
suffix := time.Now().UnixNano()
first, err := d.CreateTask(fmt.Sprintf("task-pin-first-%d", suffix), "goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
second, err := d.CreateTask(fmt.Sprintf("task-pin-second-%d", suffix), "goal", nil, 0, 0)
if err != nil {
_ = d.DeleteTask(first.ID)
t.Fatal(err)
}
defer func() {
_ = d.DeleteTask(first.ID)
_ = d.DeleteTask(second.ID)
}()
pinned := true
first, err = d.UpdateTask(first.ID, TaskPatch{Pinned: &pinned})
if err != nil || first == nil || !first.Pinned || first.PinnedAt == nil {
t.Fatalf("pin first = %+v, %v", first, err)
}
firstPinnedAt := *first.PinnedAt
time.Sleep(5 * time.Millisecond)
second, err = d.UpdateTask(second.ID, TaskPatch{Pinned: &pinned})
if err != nil || second == nil || !second.Pinned || second.PinnedAt == nil {
t.Fatalf("pin second = %+v, %v", second, err)
}
name := "renamed pinned task"
first, err = d.UpdateTask(first.ID, TaskPatch{Name: &name, Pinned: &pinned})
if err != nil || first == nil {
t.Fatal(err)
}
if first.Name != name || first.PinnedAt == nil || !first.PinnedAt.Equal(firstPinnedAt) {
t.Fatalf("repeat pin should preserve pin time: %+v (want %v)", first, firstPinnedAt)
}
tasks, err := d.ListTasks()
if err != nil {
t.Fatal(err)
}
positions := map[int64]int{}
for index, task := range tasks {
positions[task.ID] = index
}
if positions[second.ID] >= positions[first.ID] {
t.Fatalf("newer pin must sort first: second=%d first=%d", positions[second.ID], positions[first.ID])
}
pinned = false
first, err = d.UpdateTask(first.ID, TaskPatch{Pinned: &pinned})
if err != nil || first == nil || first.Pinned || first.PinnedAt != nil {
t.Fatalf("unpin first = %+v, %v", first, err)
}
}
func TestDeleteConversationsReturnsExistingIDs(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
first, err := d.CreateConversation("mainagent", "batch-delete-first", nil)
if err != nil {
t.Fatal(err)
}
second, err := d.CreateConversation("mainagent", "batch-delete-second", nil)
if err != nil {
_ = d.DeleteConversation(first.ID)
t.Fatal(err)
}
defer func() {
_ = d.DeleteConversation(first.ID)
_ = d.DeleteConversation(second.ID)
}()
missing := int64(1<<62 - 1)
deleted, err := d.DeleteConversations([]int64{first.ID, missing, second.ID})
if err != nil {
t.Fatal(err)
}
seen := make(map[int64]bool, len(deleted))
for _, id := range deleted {
seen[id] = true
}
if !seen[first.ID] || !seen[second.ID] || seen[missing] || len(deleted) != 2 {
t.Fatalf("deleted=%v, want existing ids only", deleted)
}
}
+66
View File
@@ -0,0 +1,66 @@
package db
import (
"testing"
"time"
)
func TestTaskQueuePreservesBootstrapAndFIFOPosition(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
task, err := d.CreateTask("queue metadata", "keep first-run mode", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
defer d.DeleteTask(task.ID)
if err := d.Enqueue(task.ID, "bootstrap"); err != nil {
t.Fatal(err)
}
first, err := d.GetTask(task.ID)
if err != nil || first == nil || first.QueuedAt == nil {
t.Fatalf("first enqueue: task=%+v err=%v", first, err)
}
// A follow-up/rerun can try to admit an already queued task as resume. The
// original bootstrap mode and FIFO timestamp must remain authoritative.
time.Sleep(time.Millisecond)
if err := d.Enqueue(task.ID, "resume"); err != nil {
t.Fatal(err)
}
second, err := d.GetTask(task.ID)
if err != nil || second == nil {
t.Fatalf("second enqueue: task=%+v err=%v", second, err)
}
if second.QueueMode != "bootstrap" {
t.Fatalf("queue mode=%q, want bootstrap", second.QueueMode)
}
if second.QueuedAt == nil || !second.QueuedAt.Equal(*first.QueuedAt) {
t.Fatalf("repeated enqueue moved FIFO position: first=%v second=%v", first.QueuedAt, second.QueuedAt)
}
// Pausing a queued task removes it from the queue but keeps its required
// startup mode. A later requeue receives a new tail position.
if err := d.Dequeue(task.ID, false); err != nil {
t.Fatal(err)
}
paused, err := d.GetTask(task.ID)
if err != nil || paused == nil || paused.Queued || paused.QueueMode != "bootstrap" {
t.Fatalf("paused queue metadata: task=%+v err=%v", paused, err)
}
time.Sleep(time.Millisecond)
if err := d.Enqueue(task.ID, "bootstrap"); err != nil {
t.Fatal(err)
}
requeued, err := d.GetTask(task.ID)
if err != nil || requeued == nil || requeued.QueuedAt == nil {
t.Fatalf("requeue: task=%+v err=%v", requeued, err)
}
if !requeued.QueuedAt.After(*first.QueuedAt) {
t.Fatalf("requeue did not move to FIFO tail: first=%v requeued=%v", first.QueuedAt, requeued.QueuedAt)
}
}
+644
View File
@@ -0,0 +1,644 @@
package db
import (
"database/sql"
"fmt"
"net"
"strconv"
"strings"
)
// TaskScope is one row of a task's test scope — the coverage denominator and the
// per-task authorization edge. Rows come either from insertAssets (source='auto',
// conservative, one per explicitly-inserted asset) or from the add_task_scope tool
// (source='agent', for company / domain / network / ICP / keyword scope).
type TaskScope struct {
ID int64 `json:"id"`
TaskID int64 `json:"task_id"`
Kind string `json:"kind"` // company|root_domain|subdomain|ip|cidr|icp|keyword
CompanyID *int64 `json:"company_id,omitempty"`
// CompanyName is resolved for kind=company so callers can label a scope row
// without a second lookup. Empty when the row is not a company reference.
CompanyName string `json:"company_name,omitempty"`
Domain string `json:"domain,omitempty"`
Net string `json:"net,omitempty"`
Value string `json:"value,omitempty"`
Source string `json:"source"`
Reason string `json:"reason,omitempty"`
}
// stripHostPort drops a trailing :port from a host:port / ip:port / [ipv6]:port
// value, returning the bare host. Bare hosts, bare IPs (v4 or v6, whose colons make
// them ambiguous), and anything not in host:port form are returned unchanged. Scope
// is host/net based, so the port from a target like "10.0.188.136:3000" or
// "api.example.com:8080" is simply discarded rather than baked into the key.
func stripHostPort(v string) string {
v = strings.TrimSpace(v)
if host, _, err := net.SplitHostPort(v); err == nil {
return host
}
return v
}
// ipToHostCIDR turns a bare IP into its single-host CIDR (/32 or /128). "" if invalid.
func ipToHostCIDR(ip string) string {
ip = strings.TrimSpace(ip)
p := net.ParseIP(ip)
if p == nil {
return ""
}
if p.To4() != nil {
return ip + "/32"
}
return ip + "/128"
}
// upsertTaskScope inserts one scope row idempotently (uq_task_scope). A duplicate is
// silently ignored. taskID<=0 or empty kind → no-op.
func (s *AssetStore) upsertTaskScopeResult(ts TaskScope) (bool, error) {
if ts.TaskID <= 0 || ts.Kind == "" {
return false, nil
}
var domainVal, netVal, companyVal, valueVal any
if ts.Domain != "" {
domainVal = ts.Domain
}
if ts.Net != "" {
netVal = ts.Net
}
if ts.CompanyID != nil && *ts.CompanyID > 0 {
companyVal = *ts.CompanyID
}
if ts.Value != "" {
valueVal = ts.Value
}
src := ts.Source
if src == "" {
src = "auto"
}
query := `
INSERT INTO task_scope(task_id, kind, company_id, domain, net, value, source, reason)
VALUES ($1,$2,$3,$4,$5::cidr,$6,$7,NULLIF($8,''))
ON CONFLICT DO NOTHING`
var (
result sql.Result
err error
)
if s.tx != nil {
result, err = s.tx.Exec(query, ts.TaskID, ts.Kind, companyVal, domainVal, netVal, valueVal, src, ts.Reason)
} else {
result, err = s.db.Exec(query, ts.TaskID, ts.Kind, companyVal, domainVal, netVal, valueVal, src, ts.Reason)
}
if err != nil {
return false, err
}
rows, err := result.RowsAffected()
return rows > 0, err
}
func (s *AssetStore) upsertTaskScope(ts TaskScope) error {
_, err := s.upsertTaskScopeResult(ts)
return err
}
// AddAutoScope records the conservative task scope implied by ONE explicitly-inserted
// asset item (source='auto'). MUST be called only from insertAssets' top-level loop —
// never from a db-layer side effect (linkHostAssets), so派生资产不会盲目扩大范围。
// Rule: scope granularity follows the asset's own type. taskID<=0 → no-op.
func (s *AssetStore) AddAutoScope(taskID int64, assetType, domain, rawURL, ip string) error {
if taskID <= 0 {
return nil
}
switch assetType {
case "root_domain":
if d := DomainKey(domain); d != "" {
return s.upsertTaskScope(TaskScope{TaskID: taskID, Kind: "root_domain", Domain: d})
}
case "subdomain":
if d := DomainKey(domain); d != "" {
return s.upsertTaskScope(TaskScope{TaskID: taskID, Kind: "subdomain", Domain: d})
}
case "service", "endpoint":
host := domain
if host == "" && rawURL != "" {
host, _, _ = parseURL(normalizeURL(rawURL))
}
host = DomainKey(host)
if host != "" && net.ParseIP(host) == nil {
return s.upsertTaskScope(TaskScope{TaskID: taskID, Kind: "subdomain", Domain: host})
}
// IP-literal host or no host → fall back to ip scope if we have one.
if c := ipToHostCIDR(ip); c != "" {
return s.upsertTaskScope(TaskScope{TaskID: taskID, Kind: "ip", Net: c})
}
case "ip":
if c := ipToHostCIDR(ip); c != "" {
return s.upsertTaskScope(TaskScope{TaskID: taskID, Kind: "ip", Net: c})
}
}
return nil
}
// AddAgentScope parses a (kind,value) pair and records it as task scope.
// source is "agent" when called from an LLM tool, "manual" from the UI.
// company: value = company name or id (must already exist). root_domain/subdomain:
// value = a domain. ip/cidr: value = an IP or CIDR (bare IP → /32,/128).
func (s *AssetStore) AddAgentScope(taskID int64, kind, value, reason, source string) (TaskScope, error) {
if source == "" {
source = "agent"
}
ts := TaskScope{TaskID: taskID, Kind: kind, Source: source, Reason: reason}
if taskID <= 0 {
return ts, fmt.Errorf("需要 task_id")
}
value = strings.TrimSpace(value)
if value == "" {
return ts, fmt.Errorf("value 不能为空")
}
switch kind {
case "company":
if s.company == nil {
return ts, fmt.Errorf("company store 未启用")
}
var comp *Company
var err error
if id, e := strconv.ParseInt(value, 10, 64); e == nil {
comp, err = s.company.GetCompany(id)
} else {
comp, err = s.company.GetCompanyByName(value)
}
if err != nil {
return ts, err
}
if comp == nil {
return ts, fmt.Errorf("company 不存在: %s(先用 list_companies 确认,或建好企业)", value)
}
ts.CompanyID = &comp.ID
case "root_domain":
d := DomainKey(stripHostPort(value))
root, _ := RootDomain(d)
if root == "" {
root = d
}
if root == "" {
return ts, fmt.Errorf("无效根域: %s", value)
}
ts.Domain = root
case "subdomain":
d := DomainKey(stripHostPort(value))
if d == "" {
return ts, fmt.Errorf("无效子域: %s", value)
}
ts.Domain = d
case "ip", "cidr":
v := value
if !strings.Contains(v, "/") {
v = ipToHostCIDR(stripHostPort(v))
ts.Kind = "ip"
} else {
ts.Kind = "cidr"
}
if v == "" {
return ts, fmt.Errorf("无效 ip/cidr: %s", value)
}
if _, _, err := net.ParseCIDR(v); err != nil {
return ts, fmt.Errorf("无效 ip/cidr: %s", value)
}
ts.Net = v
case "icp", "keyword":
parsed, err := ParseScopeInput(ScopeInput{Kind: kind, Value: value})
if err != nil {
return ts, err
}
ts.Value = parsed.Value
default:
return ts, fmt.Errorf("不支持的 kind: %s(company/root_domain/subdomain/ip/cidr/icp/keyword)", kind)
}
if err := s.upsertTaskScope(ts); err != nil {
return ts, err
}
return ts, nil
}
// DeleteTaskScope removes a single scope row by id, scoped to the given task.
// Returns whether a row was actually deleted.
func (s *AssetStore) DeleteTaskScope(taskID, scopeID int64) (bool, error) {
res, err := s.db.Exec(`DELETE FROM task_scope WHERE id=$1 AND task_id=$2`, scopeID, taskID)
if err != nil {
return false, err
}
n, _ := res.RowsAffected()
return n > 0, nil
}
// ListTaskScope returns all scope rows for a task.
func (s *AssetStore) ListTaskScope(taskID int64) ([]TaskScope, error) {
rows, err := s.db.Query(`
SELECT ts.id, ts.kind, COALESCE(ts.company_id,0), COALESCE(c.name,''), COALESCE(ts.domain,''),
COALESCE(ts.net::text,''), COALESCE(ts.value,''), ts.source, COALESCE(ts.reason,'')
FROM task_scope ts
LEFT JOIN companies c ON c.id=ts.company_id
WHERE ts.task_id=$1 ORDER BY ts.id`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
out := []TaskScope{}
for rows.Next() {
var t TaskScope
var cid int64
if err := rows.Scan(&t.ID, &t.Kind, &cid, &t.CompanyName, &t.Domain, &t.Net, &t.Value, &t.Source, &t.Reason); err != nil {
return nil, err
}
t.TaskID = taskID
if cid > 0 {
t.CompanyID = &cid
}
out = append(out, t)
}
return out, rows.Err()
}
// CoverageAsset is one in-scope asset (used for the untested backlog sample).
type CoverageAsset struct {
ID int64 `json:"id"`
Type string `json:"type"`
Label string `json:"label"`
}
// CoverageByType is per-asset-type coverage: total in scope vs tested.
type CoverageByType struct {
Type string `json:"type"`
Total int `json:"total"`
Tested int `json:"tested"`
}
// Coverage is a task's rough asset test coverage — a reference figure for the agent,
// NOT a precise metric. Denominator = assets matching any active task_scope row;
// Tested = those anchored to at least one fact node in the exploration.
type Coverage struct {
Enabled bool `json:"enabled"` // 资产覆盖度功能是否开启;false 时其余字段为零值
ScopeRows int `json:"scope_rows"` // 0 → 范围未锚定
Denominator int `json:"denominator"` // 范围内资产数
Tested int `json:"tested"` // 已测(约)
Pct *float64 `json:"pct"` // 覆盖度;分母 0 时 null
ByType []CoverageByType `json:"by_type"` // 按资产类型的 总数/已测
}
// CoverageEnabled reports whether a task has the asset-coverage feature turned on
// (tasks.coverage_enabled). Missing row / error → true (fail open to the default),
// so unknown/legacy tasks keep the historical behavior. taskID<=0 → true.
func (s *AssetStore) CoverageEnabled(taskID int64) bool {
if taskID <= 0 {
return true
}
var enabled bool
if err := s.db.QueryRow(`SELECT COALESCE(coverage_enabled,true) FROM tasks WHERE id=$1`, taskID).Scan(&enabled); err != nil {
return true
}
return enabled
}
// task_scope→assets match predicate, reused by the count / by-type / untested queries. $1=taskID.
const covTargetCTE = `
target AS (
SELECT DISTINCT a.id, a.type,
COALESCE(a.url, a.domain, a.ip, a.app_name, a.root_domain, '') AS label
FROM assets a
JOIN task_scope ts ON ts.task_id = $1 AND (
(ts.kind='company' AND a.company_id = ts.company_id)
OR (ts.kind='root_domain' AND a.root_domain = ts.domain)
OR (ts.kind='subdomain' AND a.domain = ts.domain)
OR (ts.kind IN ('ip','cidr') AND ts.net >>= try_inet(a.ip))
OR (ts.kind='icp' AND (
lower(regexp_replace(COALESCE(a.icp,''), '[[:space:]]+', '', 'g')) = ts.value
OR lower(regexp_replace(COALESCE(a.app_icp,''), '[[:space:]]+', '', 'g')) = ts.value
))
)
),
tested AS (
SELECT DISTINCT ea.asset_id
FROM exploration_anchors ea
JOIN exploration_nodes en ON en.id = ea.node_id
WHERE en.exploration_id = $2 AND en.kind = 'fact'
)`
// TaskCoverage computes rough per-type coverage for a task. taskID indexes
// task_scope + assets; expID indexes the fact anchors. Reference figure only.
func (s *AssetStore) TaskCoverage(taskID, expID int64) (*Coverage, error) {
cov := &Coverage{ByType: []CoverageByType{}}
_ = s.db.QueryRow(`SELECT count(*) FROM task_scope WHERE task_id=$1`, taskID).Scan(&cov.ScopeRows)
rows, err := s.db.Query(`WITH `+covTargetCTE+`
SELECT t.type, count(*) AS total,
count(*) FILTER (WHERE t.id IN (SELECT asset_id FROM tested)) AS tested
FROM target t GROUP BY t.type ORDER BY t.type`, taskID, expID)
if err != nil {
return nil, err
}
defer rows.Close()
for rows.Next() {
var bt CoverageByType
if err := rows.Scan(&bt.Type, &bt.Total, &bt.Tested); err != nil {
return nil, err
}
cov.ByType = append(cov.ByType, bt)
cov.Denominator += bt.Total
cov.Tested += bt.Tested
}
if err := rows.Err(); err != nil {
return nil, err
}
if cov.Denominator > 0 {
p := float64(cov.Tested) / float64(cov.Denominator)
cov.Pct = &p
}
return cov, nil
}
// ---------------------------------------------------------------------------
// Coverage graph — a force-directed view of a task's in-scope assets.
// ---------------------------------------------------------------------------
// CoverageGraphNode is one node of the asset coverage graph. Node identity is a
// string key: an asset row is "a:<id>", a company is "c:<id>", and a root domain
// with no asset row of its own is the synthetic "r:<domain>". Out-of-scope
// connector nodes (roots pulled in only to link subdomains, and the companies
// above them) carry InScope=false and are rendered gray.
type CoverageGraphNode struct {
Key string `json:"key"`
Kind string `json:"kind"` // company|root_domain|subdomain|ip|service|app|endpoint
Label string `json:"label"`
Tested bool `json:"tested"`
InScope bool `json:"in_scope"`
AssetID int64 `json:"asset_id,omitempty"` // 0 for company / synthetic root
CompanyID int64 `json:"company_id,omitempty"`
Domain string `json:"domain,omitempty"`
RootDomain string `json:"root_domain,omitempty"`
IP string `json:"ip,omitempty"`
URL string `json:"url,omitempty"`
Port int `json:"port,omitempty"`
ServiceType string `json:"service_type,omitempty"`
AppName string `json:"app_name,omitempty"`
PageTitle string `json:"page_title,omitempty"`
StatusCode int `json:"status_code,omitempty"`
}
// CoverageGraphEdge is a child→parent containment link (endpoint→service→
// subdomain/ip→root_domain→company, app→company).
type CoverageGraphEdge struct {
Src string `json:"src"`
Dst string `json:"dst"`
}
// CoverageGraphData is the whole graph for one task.
type CoverageGraphData struct {
Nodes []CoverageGraphNode `json:"nodes"`
Edges []CoverageGraphEdge `json:"edges"`
}
func assetKey(id int64) string { return "a:" + strconv.FormatInt(id, 10) }
func companyKey(id int64) string { return "c:" + strconv.FormatInt(id, 10) }
// hostPortOf returns the (host, port) a service / endpoint node hangs off — the
// domain if set, otherwise the URL host, otherwise the IP.
func hostPortOf(n *CoverageGraphNode) (string, int) {
host, port := n.Domain, n.Port
if host == "" && n.URL != "" {
h, p, _ := parseURL(normalizeURL(n.URL))
host = h
if port == 0 {
port = p
}
}
if host == "" {
host = n.IP
}
return host, port
}
// BuildCoverageGraph assembles the full coverage graph for a task and its direct
// read-only sources: every in-scope asset plus connector root domains/companies,
// with current-or-source fact anchors reflected in Tested. The legacy expID
// argument is retained for API compatibility; the task registry is authoritative.
func (s *AssetStore) BuildCoverageGraph(taskID, _ int64) (*CoverageGraphData, error) {
g := &CoverageGraphData{Nodes: []CoverageGraphNode{}, Edges: []CoverageGraphEdge{}}
if taskID <= 0 {
return g, nil
}
rows, err := s.db.Query(`WITH `+contextCoverageCTE+`
SELECT a.id, a.type, COALESCE(a.company_id,0),
COALESCE(a.domain,''), COALESCE(a.root_domain,''), COALESCE(a.ip,''),
COALESCE(a.url,''), COALESCE(a.port,0), COALESCE(a.service_type,''),
COALESCE(a.app_name,''), COALESCE(a.page_title,''), COALESCE(a.status_code,0),
(a.id IN (SELECT asset_id FROM tested)) AS tested
FROM assets a JOIN target t ON t.id = a.id
ORDER BY a.id`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
byKey := map[string]*CoverageGraphNode{}
rootByDomain := map[string]string{} // root domain → node key
subByDomain := map[string]string{} // subdomain → node key
ipByAddr := map[string]string{} // ip literal → node key
svcByHostPort := map[string]string{}
svcByHost := map[string]string{}
companyIDs := map[int64]bool{} // referenced company ids (need a node)
add := func(n CoverageGraphNode) *CoverageGraphNode {
if _, ok := byKey[n.Key]; ok {
return byKey[n.Key]
}
g.Nodes = append(g.Nodes, n)
p := &g.Nodes[len(g.Nodes)-1]
byKey[n.Key] = p
return p
}
for rows.Next() {
var n CoverageGraphNode
var companyID int64
if err := rows.Scan(&n.AssetID, &n.Kind, &companyID,
&n.Domain, &n.RootDomain, &n.IP, &n.URL, &n.Port, &n.ServiceType,
&n.AppName, &n.PageTitle, &n.StatusCode, &n.Tested); err != nil {
return nil, err
}
n.Key = assetKey(n.AssetID)
n.CompanyID = companyID
n.InScope = true
n.Label = coverageNodeLabel(&n)
p := add(n)
switch n.Kind {
case "root_domain":
if n.Domain != "" {
rootByDomain[n.Domain] = p.Key
}
case "subdomain":
if n.Domain != "" {
subByDomain[n.Domain] = p.Key
}
case "ip":
if n.IP != "" {
ipByAddr[n.IP] = p.Key
}
case "service":
host, port := hostPortOf(p)
if host != "" {
svcByHost[host] = p.Key
svcByHostPort[host+"|"+strconv.Itoa(port)] = p.Key
}
}
if companyID > 0 {
companyIDs[companyID] = true
}
}
if err := rows.Err(); err != nil {
return nil, err
}
// Connector root domains: any subdomain whose root domain is not itself an
// in-scope node. Pull the real asset row if one exists (so the drawer shows
// real detail), else synthesize a bare "r:<domain>" placeholder. Both gray.
missingRoots := map[string]bool{}
for _, n := range g.Nodes {
if n.Kind == "subdomain" && n.RootDomain != "" {
if _, ok := rootByDomain[n.RootDomain]; !ok {
missingRoots[n.RootDomain] = true
}
}
}
for root := range missingRoots {
var id, companyID int64
err := s.db.QueryRow(`SELECT id, COALESCE(company_id,0) FROM assets
WHERE type='root_domain' AND domain=$1 LIMIT 1`, root).Scan(&id, &companyID)
var node CoverageGraphNode
if err == nil && id > 0 {
node = CoverageGraphNode{Key: assetKey(id), Kind: "root_domain", AssetID: id,
CompanyID: companyID, Domain: root, Label: root}
if companyID > 0 {
companyIDs[companyID] = true
}
} else {
node = CoverageGraphNode{Key: "r:" + root, Kind: "root_domain", Domain: root, Label: root}
}
add(node)
rootByDomain[root] = node.Key
}
// Company nodes for every referenced company id — always gray context.
for id := range companyIDs {
key := companyKey(id)
if _, ok := byKey[key]; ok {
continue
}
var name string
if err := s.db.QueryRow(`SELECT name FROM companies WHERE id=$1`, id).Scan(&name); err != nil {
continue
}
add(CoverageGraphNode{Key: key, Kind: "company", CompanyID: id,
Label: name, AssetID: 0})
}
// Derived containment edges (only when the parent node exists).
link := func(childKey, parentKey string) {
if parentKey == "" || parentKey == childKey {
return
}
if _, ok := byKey[parentKey]; !ok {
return
}
g.Edges = append(g.Edges, CoverageGraphEdge{Src: childKey, Dst: parentKey})
}
firstOf := func(keys ...string) string {
for _, k := range keys {
if k != "" {
if _, ok := byKey[k]; ok {
return k
}
}
}
return ""
}
for i := range g.Nodes {
n := &g.Nodes[i]
switch n.Kind {
case "subdomain":
link(n.Key, rootByDomain[n.RootDomain])
case "root_domain", "app", "ip":
if n.CompanyID > 0 {
link(n.Key, companyKey(n.CompanyID))
}
case "service":
link(n.Key, firstOf(subByDomain[n.Domain], ipByAddr[n.IP], rootByDomain[n.RootDomain]))
case "endpoint":
host, port := hostPortOf(n)
parent := firstOf(
svcByHostPort[host+"|"+strconv.Itoa(port)], svcByHost[host],
subByDomain[host], subByDomain[n.Domain], ipByAddr[host], ipByAddr[n.IP],
rootByDomain[n.RootDomain])
link(n.Key, parent)
}
}
return g, nil
}
// coverageNodeLabel picks the human label for a coverage-graph node.
func coverageNodeLabel(n *CoverageGraphNode) string {
switch n.Kind {
case "endpoint", "service":
if n.URL != "" {
return n.URL
}
case "app":
if n.AppName != "" {
return n.AppName
}
}
for _, v := range []string{n.URL, n.Domain, n.IP, n.AppName, n.RootDomain} {
if v != "" {
return v
}
}
return n.Key
}
// ListUntestedAssets returns a task's in-scope, not-yet-tested assets, optionally
// filtered by asset type, paginated. Returns the page + the total count. limit<=0 → 10.
func (s *AssetStore) ListUntestedAssets(taskID, expID int64, typ string, limit, offset int) ([]CoverageAsset, int, error) {
if limit <= 0 {
limit = 10
}
if offset < 0 {
offset = 0
}
typeFilter := ""
args := []any{taskID, expID}
if typ != "" {
typeFilter = " AND t.type = $3"
args = append(args, typ)
}
var total int
_ = s.db.QueryRow(`WITH `+covTargetCTE+`
SELECT count(*) FROM target t WHERE t.id NOT IN (SELECT asset_id FROM tested)`+typeFilter, args...).Scan(&total)
pageArgs := append(append([]any{}, args...), limit, offset)
limPos := strconv.Itoa(len(args) + 1)
offPos := strconv.Itoa(len(args) + 2)
rows, err := s.db.Query(`WITH `+covTargetCTE+`
SELECT t.id, t.type, t.label FROM target t
WHERE t.id NOT IN (SELECT asset_id FROM tested)`+typeFilter+`
ORDER BY t.id LIMIT $`+limPos+` OFFSET $`+offPos, pageArgs...)
if err != nil {
return nil, total, err
}
defer rows.Close()
out := []CoverageAsset{}
for rows.Next() {
var a CoverageAsset
if err := rows.Scan(&a.ID, &a.Type, &a.Label); err != nil {
return nil, total, err
}
out = append(out, a)
}
return out, total, rows.Err()
}
+23
View File
@@ -0,0 +1,23 @@
package db
import "testing"
// TestStripHostPort covers the ip:port / [ipv6]:port stripping used by the
// ip/cidr scope path so a target like "10.0.188.136:3000" no longer 404s.
func TestStripHostPort(t *testing.T) {
cases := []struct{ in, want string }{
{"10.0.188.136:3000", "10.0.188.136"}, // the reported IP case
{"10.0.188.136", "10.0.188.136"}, // bare IPv4 unchanged
{"[2001:db8::1]:8080", "2001:db8::1"}, // bracketed IPv6 + port
{"2001:db8::1", "2001:db8::1"}, // bare IPv6 unchanged (has colons)
{" 1.2.3.4:80 ", "1.2.3.4"}, // trims surrounding space
{"example.com:443", "example.com"}, // domain + port → bare host
{"api.example.com:8080", "api.example.com"}, // subdomain + port
{"example.com", "example.com"}, // bare domain unchanged
}
for _, c := range cases {
if got := stripHostPort(c.in); got != c.want {
t.Errorf("stripHostPort(%q) = %q, want %q", c.in, got, c.want)
}
}
}
+47
View File
@@ -0,0 +1,47 @@
package db
import (
"errors"
"reflect"
"strings"
"testing"
)
func TestCreateTaskRejectsTooManySourcesBeforeOpeningTransaction(t *testing.T) {
sourceIDs := make([]int64, MaxTaskSourceCount+1)
for i := range sourceIDs {
sourceIDs[i] = int64(i + 1)
}
// No database handle is needed: validation must run before Begin so an
// oversized request cannot consume a connection or create partial rows.
_, err := (&DB{}).CreateTaskWithOptions("child", "goal", TaskCreateOptions{SourceTaskIDs: sourceIDs})
if err == nil || !strings.Contains(err.Error(), "too many source tasks") {
t.Fatalf("expected source-count validation error, got %v", err)
}
}
func TestNormalizeTaskCompanyIDs(t *testing.T) {
got, err := NormalizeTaskCompanyIDs([]int64{4, 2, 4, 7, 2})
if err != nil {
t.Fatal(err)
}
if want := []int64{4, 2, 7}; !reflect.DeepEqual(got, want) {
t.Fatalf("NormalizeTaskCompanyIDs=%v, want %v", got, want)
}
if _, err := NormalizeTaskCompanyIDs([]int64{1, 0}); !errors.Is(err, ErrTaskCompanyIDsInvalid) {
t.Fatalf("invalid company id error=%v", err)
}
}
func TestCreateTaskRejectsTooManyCompaniesBeforeOpeningTransaction(t *testing.T) {
companyIDs := make([]int64, MaxTaskCompanyCount+1)
for i := range companyIDs {
companyIDs[i] = int64(i + 1)
}
_, err := (&DB{}).CreateTaskWithOptions("child", "goal", TaskCreateOptions{CompanyIDs: companyIDs})
if !errors.Is(err, ErrTaskCompanyIDsInvalid) {
t.Fatalf("expected company-count validation error, got %v", err)
}
}
+276
View File
@@ -0,0 +1,276 @@
package db
import (
"database/sql"
"encoding/json"
"errors"
"fmt"
"strings"
"time"
"unicode/utf8"
"github.com/jackc/pgx/v5/pgconn"
)
const (
MaxTaskTemplateNameRunes = 120
MaxTaskTemplateTextRunes = 16000
)
var (
ErrTaskTemplateInvalid = errors.New("invalid task template")
ErrTaskTemplateNameConflict = errors.New("task template name already exists")
ErrTaskTemplateNotFound = errors.New("task template not found")
)
// TaskTemplate is a reusable task preset (description/goal + optional category
// and task-level intercept/allow rules).
type TaskTemplate struct {
ID int64 `json:"id"`
Name string `json:"name"`
NKey string `json:"-"`
Description string `json:"description"`
Goal string `json:"goal"`
CategoryID *int64 `json:"category_id"`
InterceptRules []TaskInterceptRuleInput `json:"intercept_rules"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// TaskTemplateInput is the create/update payload after normalization.
type TaskTemplateInput struct {
Name string
Description string
Goal string
CategoryID *int64
InterceptRules []TaskInterceptRuleInput
}
// TaskTemplatePatch changes only fields flagged as set. Name/Description/Goal use
// non-nil pointers; CategoryID/InterceptRules use explicit Set flags (so a nil
// CategoryID can mean "clear" when SetCategoryID is true).
type TaskTemplatePatch struct {
Name *string
Description *string
Goal *string
CategoryID *int64
SetCategoryID bool
InterceptRules []TaskInterceptRuleInput
SetInterceptRules bool
}
const taskTemplateCols = `id, name, nkey, description, goal, category_id, intercept_rules, created_at, updated_at`
func scanTaskTemplate(row interface{ Scan(...any) error }) (TaskTemplate, error) {
var t TaskTemplate
var rulesRaw []byte
if err := row.Scan(&t.ID, &t.Name, &t.NKey, &t.Description, &t.Goal, &t.CategoryID, &rulesRaw, &t.CreatedAt, &t.UpdatedAt); err != nil {
return t, err
}
t.InterceptRules = []TaskInterceptRuleInput{}
if len(rulesRaw) > 0 {
if err := json.Unmarshal(rulesRaw, &t.InterceptRules); err != nil {
return t, err
}
if t.InterceptRules == nil {
t.InterceptRules = []TaskInterceptRuleInput{}
}
}
return t, nil
}
// marshalTemplateRules serializes a template's rule snapshot to JSONB text,
// always producing a JSON array (never null).
func marshalTemplateRules(rules []TaskInterceptRuleInput) ([]byte, error) {
if rules == nil {
rules = []TaskInterceptRuleInput{}
}
return json.Marshal(rules)
}
// taskTemplateName normalizes display whitespace while preserving the user's case.
func taskTemplateName(name string) string { return strings.Join(strings.Fields(name), " ") }
// taskTemplateNKey is the case-insensitive identity used by the unique index.
func taskTemplateNKey(name string) string { return strings.ToLower(taskTemplateName(name)) }
func normalizeTaskTemplateInput(in TaskTemplateInput) (TaskTemplateInput, string, error) {
in.Name = taskTemplateName(in.Name)
in.Description = strings.TrimSpace(in.Description)
in.Goal = strings.TrimSpace(in.Goal)
switch {
case in.Name == "":
return in, "", fmt.Errorf("%w: name is required", ErrTaskTemplateInvalid)
case utf8.RuneCountInString(in.Name) > MaxTaskTemplateNameRunes:
return in, "", fmt.Errorf("%w: name exceeds %d characters", ErrTaskTemplateInvalid, MaxTaskTemplateNameRunes)
case in.Description == "":
return in, "", fmt.Errorf("%w: description is required", ErrTaskTemplateInvalid)
case utf8.RuneCountInString(in.Description) > MaxTaskTemplateTextRunes:
return in, "", fmt.Errorf("%w: description exceeds %d characters", ErrTaskTemplateInvalid, MaxTaskTemplateTextRunes)
case in.Goal == "":
return in, "", fmt.Errorf("%w: goal is required", ErrTaskTemplateInvalid)
case utf8.RuneCountInString(in.Goal) > MaxTaskTemplateTextRunes:
return in, "", fmt.Errorf("%w: goal exceeds %d characters", ErrTaskTemplateInvalid, MaxTaskTemplateTextRunes)
}
return in, taskTemplateNKey(in.Name), nil
}
func normalizeTaskTemplatePatch(patch TaskTemplatePatch) (TaskTemplatePatch, *string, error) {
if patch.Name == nil && patch.Description == nil && patch.Goal == nil && !patch.SetCategoryID && !patch.SetInterceptRules {
return patch, nil, fmt.Errorf("%w: no fields supplied", ErrTaskTemplateInvalid)
}
var nkey *string
if patch.Name != nil {
name := taskTemplateName(*patch.Name)
if name == "" {
return patch, nil, fmt.Errorf("%w: name is required", ErrTaskTemplateInvalid)
}
if utf8.RuneCountInString(name) > MaxTaskTemplateNameRunes {
return patch, nil, fmt.Errorf("%w: name exceeds %d characters", ErrTaskTemplateInvalid, MaxTaskTemplateNameRunes)
}
key := taskTemplateNKey(name)
patch.Name = &name
nkey = &key
}
if patch.Description != nil {
description := strings.TrimSpace(*patch.Description)
if description == "" {
return patch, nil, fmt.Errorf("%w: description is required", ErrTaskTemplateInvalid)
}
if utf8.RuneCountInString(description) > MaxTaskTemplateTextRunes {
return patch, nil, fmt.Errorf("%w: description exceeds %d characters", ErrTaskTemplateInvalid, MaxTaskTemplateTextRunes)
}
patch.Description = &description
}
if patch.Goal != nil {
goal := strings.TrimSpace(*patch.Goal)
if goal == "" {
return patch, nil, fmt.Errorf("%w: goal is required", ErrTaskTemplateInvalid)
}
if utf8.RuneCountInString(goal) > MaxTaskTemplateTextRunes {
return patch, nil, fmt.Errorf("%w: goal exceeds %d characters", ErrTaskTemplateInvalid, MaxTaskTemplateTextRunes)
}
patch.Goal = &goal
}
return patch, nkey, nil
}
func taskTemplateUniqueViolation(err error) bool {
var pgErr *pgconn.PgError
return errors.As(err, &pgErr) && pgErr.Code == "23505"
}
// CreateTaskTemplate inserts one globally reusable preset.
func (d *DB) CreateTaskTemplate(in TaskTemplateInput) (*TaskTemplate, error) {
in, nkey, err := normalizeTaskTemplateInput(in)
if err != nil {
return nil, err
}
rulesJSON, err := marshalTemplateRules(in.InterceptRules)
if err != nil {
return nil, err
}
t, err := scanTaskTemplate(d.QueryRow(`
INSERT INTO task_templates(name, nkey, description, goal, category_id, intercept_rules)
VALUES ($1,$2,$3,$4,$5,$6)
ON CONFLICT (nkey) DO NOTHING
RETURNING `+taskTemplateCols, in.Name, nkey, in.Description, in.Goal, in.CategoryID, rulesJSON))
if err == sql.ErrNoRows {
return nil, ErrTaskTemplateNameConflict
}
if err != nil {
return nil, err
}
return &t, nil
}
// ListTaskTemplates returns the most recently maintained templates first.
func (d *DB) ListTaskTemplates() ([]*TaskTemplate, error) {
rows, err := d.Query(`SELECT ` + taskTemplateCols + ` FROM task_templates ORDER BY updated_at DESC, id DESC`)
if err != nil {
return nil, err
}
defer rows.Close()
out := []*TaskTemplate{}
for rows.Next() {
t, err := scanTaskTemplate(rows)
if err != nil {
return nil, err
}
out = append(out, &t)
}
return out, rows.Err()
}
// GetTaskTemplate returns nil when id does not exist.
func (d *DB) GetTaskTemplate(id int64) (*TaskTemplate, error) {
t, err := scanTaskTemplate(d.QueryRow(`SELECT `+taskTemplateCols+` FROM task_templates WHERE id=$1`, id))
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
return &t, nil
}
// UpdateTaskTemplate replaces the editable fields of one preset.
func (d *DB) UpdateTaskTemplate(id int64, in TaskTemplateInput) (*TaskTemplate, error) {
in, _, err := normalizeTaskTemplateInput(in)
if err != nil {
return nil, err
}
return d.PatchTaskTemplate(id, TaskTemplatePatch{
Name: &in.Name, Description: &in.Description, Goal: &in.Goal,
})
}
// PatchTaskTemplate atomically changes only the supplied fields. Keeping the
// merge in one UPDATE prevents concurrent disjoint PATCH requests from losing
// each other's changes.
func (d *DB) PatchTaskTemplate(id int64, patch TaskTemplatePatch) (*TaskTemplate, error) {
patch, nkey, err := normalizeTaskTemplatePatch(patch)
if err != nil {
return nil, err
}
rulesJSON, err := marshalTemplateRules(patch.InterceptRules)
if err != nil {
return nil, err
}
t, err := scanTaskTemplate(d.QueryRow(`UPDATE task_templates
SET name=CASE WHEN $2 THEN $3::text ELSE name END,
nkey=CASE WHEN $2 THEN $4::text ELSE nkey END,
description=CASE WHEN $5 THEN $6::text ELSE description END,
goal=CASE WHEN $7 THEN $8::text ELSE goal END,
category_id=CASE WHEN $9 THEN $10::bigint ELSE category_id END,
intercept_rules=CASE WHEN $11 THEN $12::jsonb ELSE intercept_rules END
WHERE id=$1
RETURNING `+taskTemplateCols,
id,
patch.Name != nil, patch.Name, nkey,
patch.Description != nil, patch.Description,
patch.Goal != nil, patch.Goal,
patch.SetCategoryID, patch.CategoryID,
patch.SetInterceptRules, rulesJSON,
))
if err == sql.ErrNoRows {
return nil, ErrTaskTemplateNotFound
}
if taskTemplateUniqueViolation(err) {
return nil, ErrTaskTemplateNameConflict
}
if err != nil {
return nil, err
}
return &t, nil
}
// DeleteTaskTemplate deletes one preset and reports whether it existed.
func (d *DB) DeleteTaskTemplate(id int64) (bool, error) {
result, err := d.Exec(`DELETE FROM task_templates WHERE id=$1`, id)
if err != nil {
return false, err
}
n, err := result.RowsAffected()
return n > 0, err
}
+131
View File
@@ -0,0 +1,131 @@
package db
import (
"errors"
"fmt"
"strings"
"sync"
"testing"
"time"
)
func TestTaskTemplateCRUDAndNormalizedUniqueness(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
suffix := time.Now().UnixNano()
name := fmt.Sprintf("Template %d", suffix)
created, err := d.CreateTaskTemplate(TaskTemplateInput{
Name: " " + strings.ReplaceAll(name, " ", " ") + " ",
Description: " initial description ",
Goal: " initial goal ",
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _, _ = d.DeleteTaskTemplate(created.ID) })
if created.Name != name || created.Description != "initial description" || created.Goal != "initial goal" {
t.Fatalf("template was not normalized: %+v", created)
}
if _, err := d.CreateTaskTemplate(TaskTemplateInput{
Name: strings.ToUpper(name), Description: "duplicate", Goal: "duplicate",
}); !errors.Is(err, ErrTaskTemplateNameConflict) {
t.Fatalf("duplicate create error = %v, want %v", err, ErrTaskTemplateNameConflict)
}
got, err := d.GetTaskTemplate(created.ID)
if err != nil || got == nil || got.Name != name {
t.Fatalf("GetTaskTemplate = %+v, %v", got, err)
}
listed, err := d.ListTaskTemplates()
if err != nil {
t.Fatal(err)
}
found := false
for _, template := range listed {
if template.ID == created.ID {
found = true
break
}
}
if !found {
t.Fatalf("created template %d missing from list", created.ID)
}
updatedName := name + " updated"
updated, err := d.UpdateTaskTemplate(created.ID, TaskTemplateInput{
Name: updatedName, Description: "new description", Goal: "new goal",
})
if err != nil {
t.Fatal(err)
}
if updated.Name != updatedName || updated.Description != "new description" || updated.Goal != "new goal" {
t.Fatalf("unexpected updated template: %+v", updated)
}
if _, err := d.CreateTaskTemplate(TaskTemplateInput{Name: "", Description: "x", Goal: "y"}); !errors.Is(err, ErrTaskTemplateInvalid) {
t.Fatalf("empty name error = %v, want %v", err, ErrTaskTemplateInvalid)
}
deleted, err := d.DeleteTaskTemplate(created.ID)
if err != nil || !deleted {
t.Fatalf("DeleteTaskTemplate = %v, %v", deleted, err)
}
deleted, err = d.DeleteTaskTemplate(created.ID)
if err != nil || deleted {
t.Fatalf("second DeleteTaskTemplate = %v, %v", deleted, err)
}
}
func TestTaskTemplateDisjointPatchesCompose(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
created, err := d.CreateTaskTemplate(TaskTemplateInput{
Name: fmt.Sprintf("Concurrent template %d", time.Now().UnixNano()),
Description: "initial description",
Goal: "initial goal",
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _, _ = d.DeleteTaskTemplate(created.ID) })
description := "description from concurrent patch"
goal := "goal from concurrent patch"
start := make(chan struct{})
errs := make(chan error, 2)
var wg sync.WaitGroup
for _, patch := range []TaskTemplatePatch{{Description: &description}, {Goal: &goal}} {
patch := patch
wg.Add(1)
go func() {
defer wg.Done()
<-start
_, err := d.PatchTaskTemplate(created.ID, patch)
errs <- err
}()
}
close(start)
wg.Wait()
close(errs)
for err := range errs {
if err != nil {
t.Fatal(err)
}
}
got, err := d.GetTaskTemplate(created.ID)
if err != nil {
t.Fatal(err)
}
if got == nil || got.Description != description || got.Goal != goal {
t.Fatalf("disjoint patches lost an update: %+v", got)
}
}
+654
View File
@@ -0,0 +1,654 @@
package db
import (
"database/sql"
"encoding/json"
"errors"
"fmt"
"strconv"
"time"
)
// Task is a row in the task registry (1:1 with an exploration).
type Task struct {
ID int64 `json:"id"`
Name string `json:"name"` // 可选任务名称;空=未命名
CategoryID *int64 `json:"category_id,omitempty"`
CategoryName string `json:"category_name,omitempty"`
Pinned bool `json:"pinned"`
PinnedAt *time.Time `json:"pinned_at,omitempty"`
Description string `json:"description"`
Goal string `json:"goal"`
ExplorationID int64 `json:"exploration_id"`
Status string `json:"status"`
Paused bool `json:"paused"`
Queued bool `json:"queued"`
QueuedAt *time.Time `json:"queued_at,omitempty"`
QueueMode string `json:"queue_mode,omitempty"`
LLMProfileID *int64 `json:"llm_profile_id,omitempty"`
// Task-level ordered LLM chain. LLMProfileID remains the compatibility alias
// for ActiveLLMProfileID while older API clients still send one profile id.
LLMProfileIDs []int64 `json:"llm_profile_ids,omitempty"`
ActiveLLMProfileID *int64 `json:"active_llm_profile_id,omitempty"`
LLMChainRevision int64 `json:"-"`
LLMFailoverState string `json:"llm_failover_state,omitempty"`
LLMFailoverReason string `json:"llm_failover_reason,omitempty"`
SourceTaskIDs []int64 `json:"source_task_ids,omitempty"`
CompanyIDs []int64 `json:"company_ids,omitempty"`
ParentRef string `json:"parent_ref,omitempty"` // 父任务 id(编排 spawn 记录;空=顶层)
CreatedAt time.Time `json:"created_at"`
CompletedAt *time.Time `json:"completed_at,omitempty"` // 进入终态(done/failed/timeout)的时刻;非终态为 nil
// 任务级超时(见 docs/任务级超时与收尾设计.md)。
TimeoutSeconds int `json:"timeout_seconds"` // 0=不限时
FirstRunAt *time.Time `json:"first_run_at,omitempty"` // 首次真正开始运行的时刻(非 created_at);nil=尚未运行
DeadlineAt *time.Time `json:"deadline_at,omitempty"` // = first_run_at + timeout_seconds;nil=不限或未运行
// planner 心跳触发间隔(秒):距上轮 plan 结束/任务开始满该值且期间无触发 → 触发一轮。
// 下限=默认=300(5min),低于一律抬到 300(在 CreateTask 归一)。见 docs/planner-trigger-impl-plan.md
PlanHeartbeatSeconds int `json:"plan_heartbeat_seconds"`
// CoverageEnabled 是「资产覆盖度功能」总开关(默认 true)。false 时:不计算/不展示测试
// 覆盖度、不自动累积 task_scope(source=auto)、不给 agent 开放 add_task_scope/
// list_untested_assets、态势里不注入 coverage 块(scope 字段仍保留)。company 关联
// (task_scope kind=company)与此开关无关,永不受影响。见 db/task_scope.go。
CoverageEnabled bool `json:"coverage_enabled"`
}
// TaskDeleteResult reports optional related-data cleanup performed in the same
// transaction as the task/exploration delete.
type TaskDeleteResult struct {
AssetsDeleted int64
AssetsDetached int64
FindingsDeleted int64
LLMRecordsDeleted int64
}
// TaskDeletePreparation is produced inside the PostgreSQL deletion transaction
// after asset and anchor writers have been excluded. Prepare callbacks may use
// TrafficHosts to stage an external traffic deletion before PostgreSQL commits.
type TaskDeletePreparation struct {
ExplorationID int64
TrafficHosts []string
}
// IsTerminal reports whether a task status is a terminal (finished) state.
// 单一真源,替换散落各处的 done/failed 硬编码判定。
func IsTerminal(status string) bool {
return status == "done" || status == "failed" || status == "timeout"
}
// CreateTask creates an exploration + task in one transaction and returns the task.
// timeoutSeconds is the task-level wall-clock budget (0 = 不限时); deadline_at is
// stamped later at first real run (see engine), not here.
// MinPlanHeartbeatSeconds 是 planner 心跳间隔的下限 = 默认 = 10min。
// 低于它(含缺省 0 / 负值 / 误配的小值)一律抬到 10min,防止把 planner 打爆。
const MinPlanHeartbeatSeconds = 600
// MaxTaskSourceCount bounds the amount of live inherited context one task can
// pull into every planner/main-agent prompt. Inheritance is intentionally direct
// only; keeping the fan-in bounded also prevents a single create request from
// multiplying graph and asset-context queries without limit.
const MaxTaskSourceCount = 8
// MaxTaskCompanyCount bounds the number of company asset scopes attached to a
// task. Company scopes are prompt context, and their currently attributed
// assets are snapshotted into the task at creation time.
const MaxTaskCompanyCount = 32
const taskCompanyAssetSource = "company"
var (
ErrTaskCompanyIDsInvalid = errors.New("invalid task company ids")
ErrTaskCompanyNotFound = errors.New("task company not found")
)
// NormalizeTaskCompanyIDs validates IDs and removes duplicates while retaining
// the user's first-seen order.
func NormalizeTaskCompanyIDs(ids []int64) ([]int64, error) {
if len(ids) == 0 {
return nil, nil
}
seen := make(map[int64]struct{}, min(len(ids), MaxTaskCompanyCount))
normalized := make([]int64, 0, min(len(ids), MaxTaskCompanyCount))
for _, id := range ids {
if id <= 0 {
return nil, fmt.Errorf("%w: company id must be positive", ErrTaskCompanyIDsInvalid)
}
if _, exists := seen[id]; exists {
continue
}
seen[id] = struct{}{}
normalized = append(normalized, id)
if len(normalized) > MaxTaskCompanyCount {
return nil, fmt.Errorf("%w: got more than %d unique companies", ErrTaskCompanyIDsInvalid, MaxTaskCompanyCount)
}
}
return normalized, nil
}
func normalizeHeartbeat(sec int) int {
if sec < MinPlanHeartbeatSeconds {
return MinPlanHeartbeatSeconds
}
return sec
}
func (d *DB) CreateTask(description, goal string, llmProfileID *int64, timeoutSeconds, planHeartbeatSeconds int) (*Task, error) {
var ids []int64
if llmProfileID != nil {
ids = []int64{*llmProfileID}
}
return d.CreateTaskWithOptions(description, goal, TaskCreateOptions{
LLMProfileIDs: ids, TimeoutSeconds: timeoutSeconds, PlanHeartbeatSeconds: planHeartbeatSeconds,
})
}
// TaskCreateOptions contains the task data that must be committed atomically
// with the task/exploration row.
type TaskCreateOptions struct {
Name string // 可选任务名称;空=未命名
CategoryID *int64
SourceTaskIDs []int64
CompanyIDs []int64
LLMProfileIDs []int64
TimeoutSeconds int
PlanHeartbeatSeconds int
// CoverageEnabled 是「资产覆盖度功能」开关;nil=默认开(true),让不关心该开关的创建
// 路径(编排 spawn、老 API)沿用原行为。仅 web 创建任务时可显式传 false 关闭。
CoverageEnabled *bool
// InterceptRules 是任务级资产拦截规则,创建时随任务在同一事务内写入 task_intercept_rules。
InterceptRules []TaskInterceptRuleInput
}
// CreateTaskWithOptions creates an exploration, task, direct source relations,
// and the ordered task LLM chain in one transaction.
func (d *DB) CreateTaskWithOptions(description, goal string, opts TaskCreateOptions) (*Task, error) {
if len(opts.SourceTaskIDs) > MaxTaskSourceCount {
return nil, fmt.Errorf("too many source tasks: got %d, maximum is %d", len(opts.SourceTaskIDs), MaxTaskSourceCount)
}
companyIDs, err := NormalizeTaskCompanyIDs(opts.CompanyIDs)
if err != nil {
return nil, err
}
opts.CompanyIDs = companyIDs
tx, err := d.Begin()
if err != nil {
return nil, err
}
defer tx.Rollback()
var expID int64
if err := tx.QueryRow(`INSERT INTO explorations(description, goal) VALUES ($1,$2) RETURNING id`, description, goal).Scan(&expID); err != nil {
return nil, err
}
// origin fact: the exploration graph's root, a KindFact node (state='origin')
// holding the original task description. Goals/intents/findings descend from
// it, and the seeded target asset is anchored to it as lineage (the asset graph
// is global and shared, not isolated per task). Being a fact (not a special 'begin' kind) lets every
// intent uniformly connect to a fact node, including the first ones.
originPayload, _ := json.Marshal(map[string]any{
"summary": "任务起点:" + description + ";目标:" + goal,
"description": description,
"goal": goal,
})
if _, err := tx.Exec(`
INSERT INTO exploration_nodes(exploration_id, kind, payload, priority, state, origin)
VALUES ($1, 'fact', $2, 0, 'origin', 'system')`, expID, string(originPayload)); err != nil {
return nil, err
}
if opts.TimeoutSeconds < 0 {
opts.TimeoutSeconds = 0
}
opts.PlanHeartbeatSeconds = normalizeHeartbeat(opts.PlanHeartbeatSeconds)
coverageEnabled := opts.CoverageEnabled == nil || *opts.CoverageEnabled
var categoryName string
if opts.CategoryID != nil {
if *opts.CategoryID <= 0 {
return nil, fmt.Errorf("%w: category id must be positive", ErrTaskCategoryInvalid)
}
if err := tx.QueryRow(`SELECT name FROM task_categories WHERE id=$1`, *opts.CategoryID).Scan(&categoryName); err != nil {
if err == sql.ErrNoRows {
return nil, ErrTaskCategoryNotFound
}
return nil, err
}
}
var active *int64
if len(opts.LLMProfileIDs) > 0 {
id := opts.LLMProfileIDs[0]
active = &id
}
t := &Task{
Name: opts.Name, CategoryID: opts.CategoryID, CategoryName: categoryName,
Description: description, Goal: goal, ExplorationID: expID,
LLMProfileID: active, ActiveLLMProfileID: active,
LLMProfileIDs: append([]int64(nil), opts.LLMProfileIDs...),
SourceTaskIDs: append([]int64(nil), opts.SourceTaskIDs...),
CompanyIDs: append([]int64(nil), opts.CompanyIDs...),
TimeoutSeconds: opts.TimeoutSeconds, PlanHeartbeatSeconds: opts.PlanHeartbeatSeconds,
CoverageEnabled: coverageEnabled,
}
if err := tx.QueryRow(`
INSERT INTO tasks(name, category_id, description, goal, exploration_id, llm_profile_id, active_llm_profile_id, timeout_seconds, plan_heartbeat_seconds, coverage_enabled)
VALUES ($1,$2,$3,$4,$5,$6,$6,$7,$8,$9)
RETURNING id, status, paused, created_at`, opts.Name, opts.CategoryID, description, goal, expID, active, opts.TimeoutSeconds, opts.PlanHeartbeatSeconds, coverageEnabled).Scan(&t.ID, &t.Status, &t.Paused, &t.CreatedAt); err != nil {
return nil, err
}
if err := insertTaskRelations(tx, t.ID, opts.SourceTaskIDs); err != nil {
return nil, err
}
if err := insertTaskCompanies(tx, t.ID, opts.CompanyIDs); err != nil {
return nil, err
}
if err := insertTaskLLMProfiles(tx, t.ID, opts.LLMProfileIDs); err != nil {
return nil, err
}
if err := insertTaskInterceptRules(tx, t.ID, opts.InterceptRules); err != nil {
return nil, err
}
if len(opts.LLMProfileIDs) == 0 {
t.LLMFailoverState = "default"
} else {
t.LLMFailoverState = "ready"
}
return t, tx.Commit()
}
func insertTaskCompanies(tx *sql.Tx, taskID int64, companyIDs []int64) error {
if len(companyIDs) == 0 {
return nil
}
// Company scope edits rebuild assets.company_id. Serialize the creation-time
// snapshot with those edits so the task sees one committed attribution state.
if err := lockCompanyScopeMutation(tx); err != nil {
return err
}
var inserted int
err := tx.QueryRow(`
WITH requested(company_id, position) AS (
SELECT company_id, position
FROM unnest($2::bigint[]) WITH ORDINALITY AS requested(company_id, position)
), inserted AS (
INSERT INTO task_scope(task_id, kind, company_id, source, reason)
SELECT $1, 'company', companies.id, 'manual', '작업 생성 시 연결된 회사'
FROM requested
JOIN companies ON companies.id=requested.company_id
ORDER BY requested.position
RETURNING company_id
)
SELECT count(*) FROM inserted`, taskID, companyIDs).Scan(&inserted)
if err != nil {
return err
}
if inserted != len(companyIDs) {
return fmt.Errorf("%w: one or more companies do not exist", ErrTaskCompanyNotFound)
}
// Updating task_ids fires sync_task_asset_links, which first creates generic
// source rows. The provenance upsert must therefore run afterwards so the
// company name and creation-time reason remain visible to operators.
if _, err := tx.Exec(`
UPDATE assets
SET task_ids=CASE
WHEN $1=ANY(task_ids) THEN task_ids
ELSE array_append(task_ids, $1)
END
WHERE company_id=ANY($2::bigint[])`, taskID, companyIDs); err != nil {
return err
}
if _, err := tx.Exec(`
INSERT INTO task_asset_links(task_id, asset_id, source, source_summary)
SELECT $1, asset.id, $3, '작업 생성 시 연결된 회사: ' || company.name
FROM assets asset
JOIN companies company ON company.id=asset.company_id
WHERE asset.company_id=ANY($2::bigint[])
AND $1=ANY(asset.task_ids)
ON CONFLICT (task_id, asset_id) DO UPDATE
SET source=EXCLUDED.source,
source_summary=EXCLUDED.source_summary,
source_node_id=NULL`, taskID, companyIDs, taskCompanyAssetSource); err != nil {
return err
}
return nil
}
func insertTaskRelations(tx *sql.Tx, taskID int64, sourceIDs []int64) error {
seen := map[int64]bool{}
for _, sourceID := range sourceIDs {
if sourceID <= 0 || sourceID == taskID || seen[sourceID] {
return fmt.Errorf("invalid or duplicate source task id %d", sourceID)
}
seen[sourceID] = true
res, err := tx.Exec(`INSERT INTO task_relations(task_id, source_task_id)
SELECT $1, id FROM tasks WHERE id=$2 AND deleted_at IS NULL`, taskID, sourceID)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n != 1 {
return fmt.Errorf("source task %d not found", sourceID)
}
}
return nil
}
func insertTaskLLMProfiles(tx *sql.Tx, taskID int64, profileIDs []int64) error {
seen := map[int64]bool{}
for position, profileID := range profileIDs {
if profileID <= 0 || seen[profileID] {
return fmt.Errorf("invalid or duplicate LLM profile id %d", profileID)
}
seen[profileID] = true
if _, err := tx.Exec(`INSERT INTO task_llm_profiles(task_id, profile_id, position) VALUES ($1,$2,$3)`, taskID, profileID, position); err != nil {
return err
}
}
return nil
}
const taskCols = `id, COALESCE(name,''), category_id,
COALESCE((SELECT category.name FROM task_categories category WHERE category.id=tasks.category_id),''),
description, goal, exploration_id, status, paused, queued, queued_at, COALESCE(queue_mode,''), llm_profile_id, active_llm_profile_id, COALESCE(parent_ref,''), pinned_at, created_at, completed_at, COALESCE(timeout_seconds,0), COALESCE(plan_heartbeat_seconds,300), COALESCE(coverage_enabled,true), first_run_at, deadline_at`
func scanTask(sc interface{ Scan(...any) error }) (*Task, error) {
var t Task
if err := sc.Scan(&t.ID, &t.Name, &t.CategoryID, &t.CategoryName, &t.Description, &t.Goal, &t.ExplorationID, &t.Status, &t.Paused, &t.Queued, &t.QueuedAt, &t.QueueMode, &t.LLMProfileID, &t.ActiveLLMProfileID, &t.ParentRef, &t.PinnedAt, &t.CreatedAt, &t.CompletedAt, &t.TimeoutSeconds, &t.PlanHeartbeatSeconds, &t.CoverageEnabled, &t.FirstRunAt, &t.DeadlineAt); err != nil {
return nil, err
}
t.Pinned = t.PinnedAt != nil
return &t, nil
}
// SetParentRef records a task's parent task id (编排 agent spawn_task 关联).
func (d *DB) SetParentRef(id int64, parentRef string) error {
_, err := d.Exec(`UPDATE tasks SET parent_ref=NULLIF($2,'') WHERE id=$1`, id, parentRef)
return err
}
// ListTasks returns alive tasks with pinned tasks first, then newest ids.
func (d *DB) ListTasks() ([]*Task, error) {
// id 是 BIGSERIAL,同一时刻创建的任务也有稳定且唯一的顺序。
rows, err := d.Query(`SELECT ` + taskCols + ` FROM tasks WHERE deleted_at IS NULL
ORDER BY (pinned_at IS NOT NULL) DESC, pinned_at DESC NULLS LAST, id DESC`)
if err != nil {
return nil, err
}
defer rows.Close()
var out []*Task
for rows.Next() {
t, err := scanTask(rows)
if err != nil {
return nil, err
}
out = append(out, t)
}
if err := rows.Err(); err != nil {
return nil, err
}
rows.Close()
if err := d.hydrateTasksContext(out); err != nil {
return nil, err
}
return out, nil
}
// TaskPatch updates task list metadata without changing execution state.
type TaskPatch struct {
Name *string
Pinned *bool
}
// UpdateTask applies a partial task name/pin mutation and returns the updated row.
// Re-pinning an already pinned task preserves its original position.
func (d *DB) UpdateTask(id int64, patch TaskPatch) (*Task, error) {
task, err := scanTask(d.QueryRow(`UPDATE tasks SET
name = CASE WHEN $2::boolean THEN $3 ELSE name END,
pinned_at = CASE
WHEN $4::boolean IS NULL THEN pinned_at
WHEN $4::boolean THEN COALESCE(pinned_at, now())
ELSE NULL
END
WHERE id=$1 AND deleted_at IS NULL
RETURNING `+taskCols, id, patch.Name != nil, patch.Name, patch.Pinned))
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
return task, nil
}
// GetTask returns one alive task (nil if not found/deleted).
func (d *DB) GetTask(id int64) (*Task, error) {
t, err := scanTask(d.QueryRow(`SELECT `+taskCols+` FROM tasks WHERE id=$1 AND deleted_at IS NULL`, id))
if err != nil {
if err.Error() == "sql: no rows in result set" {
return nil, nil
}
return nil, err
}
if err := d.hydrateTaskContext(t); err != nil {
return nil, err
}
return t, nil
}
// SetPaused persists a task's paused flag.
func (d *DB) SetPaused(id int64, paused bool) error {
_, err := d.Exec(`UPDATE tasks SET paused=$1 WHERE id=$2`, paused, id)
return err
}
// Enqueue places a task at the tail of the persistent admission queue. Repeating
// the operation while it is already queued keeps its original FIFO position.
func (d *DB) Enqueue(id int64, mode string) error {
if mode != "bootstrap" && mode != "resume" {
return fmt.Errorf("invalid queue mode %q", mode)
}
_, err := d.Exec(`UPDATE tasks
SET queued=true,
queued_at=CASE WHEN queued THEN COALESCE(queued_at, now()) ELSE now() END,
queue_mode=CASE
WHEN queue_mode='bootstrap' OR $2='bootstrap' THEN 'bootstrap'
ELSE 'resume'
END
WHERE id=$1`, id, mode)
return err
}
// Dequeue removes the concurrency hold. clearMode is false when a user pauses a
// queued task, preserving whether its next admission must bootstrap or resume.
func (d *DB) Dequeue(id int64, clearMode bool) error {
_, err := d.Exec(`UPDATE tasks
SET queued=false,
queued_at=NULL,
queue_mode=CASE WHEN $2 THEN '' ELSE queue_mode END
WHERE id=$1`, id, clearMode)
return err
}
// SetQueued is the compatibility helper used by older callers and tests. New
// scheduling code should use Enqueue/Dequeue so FIFO metadata is explicit.
func (d *DB) SetQueued(id int64, queued bool) error {
if queued {
return d.Enqueue(id, "bootstrap")
}
return d.Dequeue(id, true)
}
// SetStatus updates a task's lifecycle status. Entering a terminal state
// (done/failed/timeout) stamps completed_at once (COALESCE keeps the first stamp
// stable); moving back to a non-terminal state clears it, so a re-run has no stale
// finish time.
func (d *DB) SetStatus(id int64, status string) error {
_, err := d.Exec(`
UPDATE tasks
SET status = $1,
completed_at = CASE WHEN $1 IN ('done','failed','timeout') THEN COALESCE(completed_at, now()) ELSE NULL END
WHERE id = $2`, status, id)
return err
}
// SetTerminalStatusGuarded sets a terminal status only when the task is NOT already
// terminal, so a completed↔timeout race resolves to the first writer (won=true).
// Returns won=false (no error) when another terminal status already stuck — the
// caller then leaves it alone. Non-terminal transitions / re-run still use SetStatus.
func (d *DB) SetTerminalStatusGuarded(id int64, status string) (won bool, err error) {
res, err := d.Exec(`
UPDATE tasks
SET status = $1,
completed_at = COALESCE(completed_at, now())
WHERE id = $2 AND status NOT IN ('done','failed','timeout')`, status, id)
if err != nil {
return false, err
}
n, _ := res.RowsAffected()
return n > 0, nil
}
// StampFirstRun records a task's first-real-run moment and computes its absolute
// deadline (= now + timeoutSeconds). Idempotent: only stamps when first_run_at is
// still NULL, so restarts / re-entries keep the original clock. timeoutSeconds<=0
// leaves deadline_at NULL (不限时). Returns the resulting deadline (nil = 不限/未变).
func (d *DB) StampFirstRun(id int64, timeoutSeconds int) (*time.Time, error) {
var deadline *time.Time
err := d.QueryRow(`
UPDATE tasks
SET first_run_at = COALESCE(first_run_at, now()),
deadline_at = CASE
WHEN first_run_at IS NOT NULL THEN deadline_at -- 已盖过章:不动
WHEN $2 > 0 THEN now() + make_interval(secs => $2)
ELSE NULL END
WHERE id = $1
RETURNING deadline_at`, id, timeoutSeconds).Scan(&deadline)
return deadline, err
}
// DeleteTask preserves the historical behavior: remove the task and its
// exploration subgraph while retaining global assets.
func (d *DB) DeleteTask(id int64) error {
_, err := d.DeleteTaskCascade(id, false, false)
return err
}
// DeleteTaskCascade hard-deletes a task and its exploration subgraph
// (nodes/edges/activity via ON DELETE CASCADE). Optional standalone findings and
// assets owned only by this task are deleted in the same transaction; shared
// assets are retained and only have this task id detached. deleteLLMRecords is
// optional to preserve callers of the original two-option API.
func (d *DB) DeleteTaskCascade(id int64, deleteAssets, deleteFindings bool, deleteLLMRecords ...bool) (TaskDeleteResult, error) {
deleteRecords := len(deleteLLMRecords) > 0 && deleteLLMRecords[0]
return d.DeleteTaskCascadePrepared(id, deleteAssets, deleteFindings, deleteRecords, nil)
}
// DeleteTaskCascadePrepared coordinates reversible external deletion with the
// PostgreSQL cascade. When prepare is non-nil, host ownership is resolved and
// prepare is invoked inside this transaction while asset and anchor writers are
// excluded. The locks remain held through commit, closing the window where a
// host could become shared after its traffic had already been staged.
//
// prepare must only stage reversible work. Any returned error or later database
// error rolls PostgreSQL back; the caller remains responsible for rolling back
// external work that its callback staged successfully.
func (d *DB) DeleteTaskCascadePrepared(
id int64,
deleteAssets, deleteFindings, deleteLLMRecords bool,
prepare func(TaskDeletePreparation) error,
) (TaskDeleteResult, error) {
var result TaskDeleteResult
tx, err := d.Begin()
if err != nil {
return result, err
}
defer tx.Rollback()
var expID int64
if err := tx.QueryRow(`SELECT exploration_id FROM tasks WHERE id=$1 FOR UPDATE`, id).Scan(&expID); err != nil {
if err == sql.ErrNoRows {
return result, nil // already gone
}
return result, err
}
// SHARE ROW EXCLUSIVE conflicts with every INSERT/UPDATE/DELETE on these
// tables and with another coordinated deletion. Taking both in one fixed order
// prevents phantoms (a distinct asset row for the same host) as well as new
// ownership/anchor references until the deletion transaction commits.
if deleteAssets || prepare != nil {
if _, err := tx.Exec(`LOCK TABLE assets, exploration_anchors IN SHARE ROW EXCLUSIVE MODE`); err != nil {
return result, err
}
}
if prepare != nil {
hosts, err := hostsForTaskDeletion(tx, id, expID)
if err != nil {
return result, err
}
if err := prepare(TaskDeletePreparation{
ExplorationID: expID,
TrafficHosts: hosts,
}); err != nil {
return result, err
}
}
if deleteAssets {
// Ownership is the union of explicit task_ids and this exploration's anchors.
// An asset is deletable only when no other live task references it through
// either mechanism. This covers legacy anchor-only seeds without destroying
// evidence anchored by another task.
res, err := tx.Exec(`
WITH candidate_assets AS (
SELECT id FROM assets WHERE $1 = ANY(task_ids)
UNION
SELECT ea.asset_id
FROM exploration_anchors ea
JOIN exploration_nodes n ON n.id=ea.node_id
WHERE n.exploration_id=$2
),
deletable AS (
SELECT a.id
FROM assets a
JOIN candidate_assets c ON c.id=a.id
WHERE NOT EXISTS (
SELECT 1 FROM tasks t
WHERE t.id<>$1 AND t.deleted_at IS NULL AND t.id=ANY(a.task_ids)
) AND NOT EXISTS (
SELECT 1
FROM exploration_anchors ea
JOIN exploration_nodes n ON n.id=ea.node_id
JOIN tasks t ON t.exploration_id=n.exploration_id
WHERE ea.asset_id=a.id AND t.id<>$1 AND t.deleted_at IS NULL
)
)
DELETE FROM assets a USING deletable d WHERE a.id=d.id`, id, expID)
if err != nil {
return result, err
}
result.AssetsDeleted, _ = res.RowsAffected()
res, err = tx.Exec(`UPDATE assets SET task_ids = array_remove(task_ids, $1) WHERE $1 = ANY(task_ids)`, id)
if err != nil {
return result, err
}
result.AssetsDetached, _ = res.RowsAffected()
}
if deleteFindings {
res, err := tx.Exec(`DELETE FROM findings WHERE task_id = $1`, id)
if err != nil {
return result, err
}
result.FindingsDeleted, _ = res.RowsAffected()
}
if deleteLLMRecords {
res, err := tx.Exec(`DELETE FROM llm_records WHERE COALESCE(task_id,'')=$1`, strconv.FormatInt(id, 10))
if err != nil {
return result, err
}
result.LLMRecordsDeleted, _ = res.RowsAffected()
}
// llm_usage (the token metering ledger) is intentionally NOT deleted with the
// task — it is kept as historical accounting even after the task is gone.
if _, err := tx.Exec(`DELETE FROM tasks WHERE id=$1`, id); err != nil {
return result, err
}
if _, err := tx.Exec(`DELETE FROM explorations WHERE id=$1`, expID); err != nil {
return result, err
}
return result, tx.Commit()
}
+708
View File
@@ -0,0 +1,708 @@
package db
import (
"errors"
"fmt"
"sync"
"testing"
"time"
)
func TestTaskLifecycleAndDeleteCascade(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
tk, err := d.CreateTask("迁移测试", "目标X", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
// populate the exploration subgraph
es := d.Exploration(tk.ExplorationID)
if _, err := es.AddIntent(map[string]any{"summary": "x"}, 5, nil, "planner"); err != nil {
t.Fatal(err)
}
// pause + status
if err := d.SetPaused(tk.ID, true); err != nil {
t.Fatal(err)
}
got, err := d.GetTask(tk.ID)
if err != nil || got == nil || !got.Paused {
t.Fatalf("paused not persisted: %+v err=%v", got, err)
}
if got.Queued {
t.Fatalf("new task should not be queued: %+v", got)
}
// queued (concurrency-hold) flag round-trips independently of paused
if err := d.SetQueued(tk.ID, true); err != nil {
t.Fatal(err)
}
if g, _ := d.GetTask(tk.ID); g == nil || !g.Queued {
t.Fatalf("queued not persisted: %+v", g)
}
if err := d.SetQueued(tk.ID, false); err != nil {
t.Fatal(err)
}
if g, _ := d.GetTask(tk.ID); g == nil || g.Queued {
t.Fatalf("queued not cleared: %+v", g)
}
// list contains it
list, _ := d.ListTasks()
found := false
for _, x := range list {
if x.ID == tk.ID {
found = true
}
}
if !found {
t.Fatalf("task not in list")
}
// delete cascades exploration subgraph
if err := d.DeleteTask(tk.ID); err != nil {
t.Fatal(err)
}
if g, _ := d.GetTask(tk.ID); g != nil {
t.Fatalf("task should be gone")
}
var nodes int
d.QueryRow(`SELECT count(*) FROM exploration_nodes WHERE exploration_id=$1`, tk.ExplorationID).Scan(&nodes)
if nodes != 0 {
t.Fatalf("exploration nodes should be cascade-deleted, got %d", nodes)
}
var exps int
d.QueryRow(`SELECT count(*) FROM explorations WHERE id=$1`, tk.ExplorationID).Scan(&exps)
if exps != 0 {
t.Fatalf("exploration should be deleted, got %d", exps)
}
}
func TestTaskDeleteCascadeAssets(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
first, err := d.CreateTask("级联删除测试", "目标A", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
second, err := d.CreateTask("共享资产保留测试", "目标B", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
assets := d.Assets()
exclusiveID, err := assets.UpsertRootDomain(UpsertRootDomainReq{
Domain: fmt.Sprintf("delete-%d.example.test", first.ID),
TaskID: first.ID,
})
if err != nil {
t.Fatal(err)
}
sharedDomain := fmt.Sprintf("shared-%d.example.test", first.ID)
sharedID, err := assets.UpsertRootDomain(UpsertRootDomainReq{Domain: sharedDomain, TaskID: first.ID})
if err != nil {
t.Fatal(err)
}
if _, err := assets.UpsertRootDomain(UpsertRootDomainReq{Domain: sharedDomain, TaskID: second.ID}); err != nil {
t.Fatal(err)
}
anchorOnlyDomain := fmt.Sprintf("anchor-only-%d.example.test", first.ID)
anchorOnlyID, err := assets.UpsertRootDomain(UpsertRootDomainReq{Domain: anchorOnlyDomain})
if err != nil {
t.Fatal(err)
}
firstOrigin, err := d.Exploration(first.ExplorationID).OriginFactID()
if err != nil || firstOrigin == 0 {
t.Fatalf("first origin: id=%d err=%v", firstOrigin, err)
}
if err := d.Exploration(first.ExplorationID).Anchor(firstOrigin, anchorOnlyID); err != nil {
t.Fatal(err)
}
protectedDomain := fmt.Sprintf("other-anchor-%d.example.test", first.ID)
protectedID, err := assets.UpsertRootDomain(UpsertRootDomainReq{Domain: protectedDomain, TaskID: first.ID})
if err != nil {
t.Fatal(err)
}
secondOrigin, err := d.Exploration(second.ExplorationID).OriginFactID()
if err != nil || secondOrigin == 0 {
t.Fatalf("second origin: id=%d err=%v", secondOrigin, err)
}
if err := d.Exploration(second.ExplorationID).Anchor(secondOrigin, protectedID); err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
_ = d.DeleteTask(first.ID)
_ = d.DeleteTask(second.ID)
_, _ = assets.DeleteByIDs([]int64{exclusiveID, sharedID, anchorOnlyID, protectedID})
})
hosts, err := assets.HostsByTask(first.ID)
if err != nil {
t.Fatal(err)
}
hostSet := make(map[string]bool, len(hosts))
for _, host := range hosts {
hostSet[host] = true
}
if !hostSet[fmt.Sprintf("delete-%d.example.test", first.ID)] || !hostSet[sharedDomain] {
t.Fatalf("task hosts missing cascade fixtures: %v", hosts)
}
deletableHosts, err := assets.HostsForTaskDeletion(first.ID, first.ExplorationID)
if err != nil {
t.Fatal(err)
}
deletableSet := make(map[string]bool, len(deletableHosts))
for _, host := range deletableHosts {
deletableSet[host] = true
}
if !deletableSet[fmt.Sprintf("delete-%d.example.test", first.ID)] || !deletableSet[anchorOnlyDomain] {
t.Fatalf("exclusive or anchor-only cleanup host missing: %v", deletableHosts)
}
if deletableSet[sharedDomain] || deletableSet[protectedDomain] {
t.Fatalf("shared host was not protected from traffic cleanup: %v", deletableHosts)
}
findingID, err := d.AddFinding(first.ID, 0, "__task_delete_cascade__", "", SeverityHigh, "summary", "evidence", "tester", []int64{exclusiveID})
if err != nil {
t.Fatal(err)
}
result, err := d.DeleteTaskCascade(first.ID, true, true)
if err != nil {
t.Fatal(err)
}
if result.AssetsDeleted < 2 || result.AssetsDetached < 2 {
t.Fatalf("unexpected asset cleanup result: %+v", result)
}
if result.FindingsDeleted != 1 {
t.Fatalf("unexpected finding cleanup result: %+v", result)
}
if finding, err := d.GetFinding(findingID); err != nil || finding != nil {
t.Fatalf("finding should be deleted, got finding=%+v err=%v", finding, err)
}
remaining, err := assets.GetByIDs([]int64{exclusiveID, sharedID, anchorOnlyID, protectedID})
if err != nil {
t.Fatal(err)
}
if len(remaining) != 2 {
t.Fatalf("only shared and other-anchored assets should remain, got %+v", remaining)
}
byID := make(map[int64]*Asset, len(remaining))
for _, asset := range remaining {
byID[asset.ID] = asset
}
if shared := byID[sharedID]; shared == nil || len(shared.TaskIDs) != 1 || shared.TaskIDs[0] != second.ID {
t.Fatalf("deleted task should be detached from shared asset: %+v", shared)
}
if protected := byID[protectedID]; protected == nil || len(protected.TaskIDs) != 0 {
t.Fatalf("other task's anchor should preserve the asset after detaching ownership: %+v", protected)
}
}
func TestTaskRelationsAndLLMFailoverChain(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
suffix := time.Now().UnixNano()
profileIDs := make([]int64, 0, 3)
for i := 0; i < 3; i++ {
id, err := d.SaveProfile(&LLMProfile{
Name: fmt.Sprintf("task-chain-%d-%d", suffix, i), Format: "openai",
Model: fmt.Sprintf("model-%d", i), APIKey: "test-key",
})
if err != nil {
t.Fatal(err)
}
profileIDs = append(profileIDs, id)
}
sourceA, err := d.CreateTask("source A", "goal A", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
sourceB, err := d.CreateTask("source B", "goal B", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
child, err := d.CreateTaskWithOptions("child", "new goal", TaskCreateOptions{
SourceTaskIDs: []int64{sourceA.ID, sourceB.ID},
LLMProfileIDs: profileIDs,
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
_ = d.DeleteTask(child.ID)
_ = d.DeleteTask(sourceA.ID)
_ = d.DeleteTask(sourceB.ID)
for _, id := range profileIDs {
_ = d.DeleteProfile(id)
}
})
got, err := d.GetTask(child.ID)
if err != nil {
t.Fatal(err)
}
if fmt.Sprint(got.SourceTaskIDs) != fmt.Sprint([]int64{sourceA.ID, sourceB.ID}) {
t.Fatalf("unexpected direct sources: %v", got.SourceTaskIDs)
}
if fmt.Sprint(got.LLMProfileIDs) != fmt.Sprint(profileIDs) || got.ActiveLLMProfileID == nil || *got.ActiveLLMProfileID != profileIDs[0] {
t.Fatalf("unexpected initial chain: %+v", got)
}
initialRevision := got.LLMChainRevision
if err := d.ReplaceTaskLLMProfiles(child.ID, profileIDs, profileIDs[0]); err != nil {
t.Fatal(err)
}
staleRevision, err := d.MarkTaskLLMProfileQuotaExhaustedAtRevision(child.ID, profileIDs[0], initialRevision, "late quota from prior generation")
if err != nil {
t.Fatal(err)
}
if !staleRevision.Stale || staleRevision.Advanced || staleRevision.ChainExhausted {
t.Fatalf("old chain revision changed replacement chain: %+v", staleRevision)
}
got, err = d.GetTask(child.ID)
if err != nil {
t.Fatal(err)
}
if got.LLMChainRevision <= initialRevision || got.ActiveLLMProfileID == nil || *got.ActiveLLMProfileID != profileIDs[0] {
t.Fatalf("replacement revision/cursor not preserved: %+v", got)
}
replacementRevision := got.LLMChainRevision
transition, err := d.MarkTaskLLMProfileQuotaExhausted(child.ID, profileIDs[0], "insufficient_quota")
if err != nil {
t.Fatal(err)
}
if !transition.Advanced || transition.ChainExhausted || transition.NextProfileID == nil || *transition.NextProfileID != profileIDs[1] {
t.Fatalf("unexpected first transition: %+v", transition)
}
got, err = d.GetTask(child.ID)
if err != nil {
t.Fatal(err)
}
if got.LLMChainRevision != replacementRevision+1 {
t.Fatalf("automatic cursor advance did not increment revision: before=%d after=%d", replacementRevision, got.LLMChainRevision)
}
late, err := d.MarkTaskLLMProfileQuotaExhausted(child.ID, profileIDs[0], "late duplicate")
if err != nil {
t.Fatal(err)
}
if late.Advanced || late.NextProfileID == nil || *late.NextProfileID != profileIDs[1] {
t.Fatalf("late failure advanced past current profile: %+v", late)
}
if _, err := d.MarkTaskLLMProfileQuotaExhausted(child.ID, profileIDs[1], "quota_exceeded"); err != nil {
t.Fatal(err)
}
last, err := d.MarkTaskLLMProfileQuotaExhausted(child.ID, profileIDs[2], "余额不足")
if err != nil {
t.Fatal(err)
}
if !last.Advanced || !last.ChainExhausted || last.NextProfileID != nil {
t.Fatalf("unexpected exhausted transition: %+v", last)
}
duplicateLast, err := d.MarkTaskLLMProfileQuotaExhausted(child.ID, profileIDs[2], "late final duplicate")
if err != nil {
t.Fatal(err)
}
if duplicateLast.Advanced || !duplicateLast.ChainExhausted || duplicateLast.NextProfileID != nil {
t.Fatalf("duplicate final failure must be idempotent: %+v", duplicateLast)
}
got, err = d.GetTask(child.ID)
if err != nil {
t.Fatal(err)
}
if got.ActiveLLMProfileID != nil || got.LLMFailoverState != "chain_exhausted" {
t.Fatalf("chain exhaustion was not persisted: %+v", got)
}
// Two requests can observe the same last active profile. Only the transaction
// that clears the cursor may report an advance; the late one is idempotent.
if err := d.ReplaceTaskLLMProfiles(child.ID, []int64{profileIDs[2]}, profileIDs[2]); err != nil {
t.Fatal(err)
}
transitions := make(chan TaskLLMTransition, 2)
errs := make(chan error, 2)
var wg sync.WaitGroup
for i := 0; i < 2; i++ {
wg.Add(1)
go func() {
defer wg.Done()
transition, markErr := d.MarkTaskLLMProfileQuotaExhausted(child.ID, profileIDs[2], "concurrent final quota")
if markErr != nil {
errs <- markErr
return
}
transitions <- transition
}()
}
wg.Wait()
close(errs)
close(transitions)
for markErr := range errs {
t.Fatal(markErr)
}
advanced := 0
for transition := range transitions {
if transition.Advanced {
advanced++
}
if !transition.ChainExhausted {
t.Fatalf("concurrent final transition must report exhausted chain: %+v", transition)
}
}
if advanced != 1 {
t.Fatalf("last profile advanced %d times, want exactly once", advanced)
}
// Starting manually from the middle consumes only candidates after that cursor.
// Earlier ready profiles must not be revived by hydration or unrelated deletes.
if err := d.ReplaceTaskLLMProfiles(child.ID, []int64{profileIDs[2], profileIDs[0], profileIDs[1]}, profileIDs[0]); err != nil {
t.Fatal(err)
}
if step, err := d.MarkTaskLLMProfileQuotaExhausted(child.ID, profileIDs[0], "middle quota"); err != nil || step.NextProfileID == nil || *step.NextProfileID != profileIDs[1] {
t.Fatalf("middle cursor did not advance to its successor: transition=%+v err=%v", step, err)
}
if end, err := d.MarkTaskLLMProfileQuotaExhausted(child.ID, profileIDs[1], "tail quota"); err != nil || !end.ChainExhausted {
t.Fatalf("tail did not exhaust manual chain: transition=%+v err=%v", end, err)
}
got, err = d.GetTask(child.ID)
if err != nil {
t.Fatal(err)
}
if got.ActiveLLMProfileID != nil || got.LLMFailoverState != "chain_exhausted" {
t.Fatalf("hydration revived a profile before the manual cursor: %+v", got)
}
unrelatedID, err := d.SaveProfile(&LLMProfile{
Name: fmt.Sprintf("task-chain-unrelated-%d", suffix), Format: "openai", Model: "other", APIKey: "test-key",
})
if err != nil {
t.Fatal(err)
}
if err := d.DeleteProfile(unrelatedID); err != nil {
t.Fatal(err)
}
got, err = d.GetTask(child.ID)
if err != nil {
t.Fatal(err)
}
if got.ActiveLLMProfileID != nil || got.LLMFailoverState != "chain_exhausted" {
t.Fatalf("deleting an unrelated profile revived an exhausted chain: %+v", got)
}
// A call that finishes after its profile was removed from the chain is stale:
// preserve the new cursor and surface the original provider error upstream.
if err := d.ReplaceTaskLLMProfiles(child.ID, []int64{profileIDs[2], profileIDs[1]}, profileIDs[2]); err != nil {
t.Fatal(err)
}
stale, err := d.MarkTaskLLMProfileQuotaExhausted(child.ID, profileIDs[0], "obsolete in-flight quota")
if err != nil {
t.Fatalf("removed in-flight profile returned an internal error: %v", err)
}
if stale.Advanced || stale.ChainExhausted || stale.NextProfileID != nil {
t.Fatalf("removed in-flight profile changed the replacement chain: %+v", stale)
}
got, err = d.GetTask(child.ID)
if err != nil {
t.Fatal(err)
}
if got.ActiveLLMProfileID == nil || *got.ActiveLLMProfileID != profileIDs[2] {
t.Fatalf("stale failure changed active replacement profile: %+v", got)
}
if err := d.ReplaceTaskLLMProfiles(child.ID, []int64{profileIDs[2], profileIDs[0], profileIDs[1]}, profileIDs[0]); err != nil {
t.Fatal(err)
}
got, err = d.GetTask(child.ID)
if err != nil {
t.Fatal(err)
}
if got.ActiveLLMProfileID == nil || *got.ActiveLLMProfileID != profileIDs[0] || got.LLMFailoverState != "ready" || got.LLMFailoverReason != "" {
t.Fatalf("chain reset did not clear failure state: %+v", got)
}
if err := d.DeleteProfile(profileIDs[0]); err != nil {
t.Fatal(err)
}
got, err = d.GetTask(child.ID)
if err != nil {
t.Fatal(err)
}
if got.ActiveLLMProfileID == nil || *got.ActiveLLMProfileID != profileIDs[1] {
t.Fatalf("deleting active profile did not select next ready entry: %+v", got)
}
if err := d.DeleteProfile(profileIDs[1]); err != nil {
t.Fatal(err)
}
got, err = d.GetTask(child.ID)
if err != nil {
t.Fatal(err)
}
if got.ActiveLLMProfileID != nil || len(got.LLMProfileIDs) != 0 || got.LLMFailoverState != "default" {
t.Fatalf("deleting active chain tail must fall back instead of wrapping: %+v", got)
}
if err := d.DeleteTask(sourceA.ID); err != nil {
t.Fatal(err)
}
got, err = d.GetTask(child.ID)
if err != nil {
t.Fatal(err)
}
if fmt.Sprint(got.SourceTaskIDs) != fmt.Sprint([]int64{sourceB.ID}) {
t.Fatalf("source delete did not cascade only its relation: %v", got.SourceTaskIDs)
}
}
func TestTaskContextRejectsDuplicatesAndAllowsTerminalLLMEdits(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
source, err := d.CreateTask("source", "goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(source.ID) })
if _, err := d.CreateTaskWithOptions("bad", "goal", TaskCreateOptions{SourceTaskIDs: []int64{source.ID, source.ID}}); err == nil {
t.Fatal("duplicate source task ids should be rejected")
}
// 终态任务仍然可以改 LLM 配置链:任务结束后主 Agent 对话继续走这条链,
// 链上模型不可用时必须还能换。
profileID, err := d.SaveProfile(&LLMProfile{
Name: fmt.Sprintf("terminal-chain-%d", time.Now().UnixNano()), Format: "openai",
Model: "terminal-model", APIKey: "test-key",
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteProfile(profileID) })
task, err := d.CreateTask("terminal", "goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(task.ID) })
if err := d.SetStatus(task.ID, "done"); err != nil {
t.Fatal(err)
}
if err := d.ReplaceTaskLLMProfiles(task.ID, []int64{profileID}, profileID); err != nil {
t.Fatalf("terminal task LLM edit should be allowed: %v", err)
}
got, err := d.GetTask(task.ID)
if err != nil {
t.Fatal(err)
}
if fmt.Sprint(got.LLMProfileIDs) != fmt.Sprint([]int64{profileID}) {
t.Fatalf("terminal task chain not persisted: %v", got.LLMProfileIDs)
}
if got.ActiveLLMProfileID == nil || *got.ActiveLLMProfileID != profileID {
t.Fatalf("terminal task active profile not persisted: %v", got.ActiveLLMProfileID)
}
if err := d.ReplaceTaskLLMProfiles(task.ID, nil, 0); err != nil {
t.Fatalf("clearing a terminal task's chain should be allowed: %v", err)
}
}
func TestCreateTaskWithCompanyScopes(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
suffix := time.Now().UnixNano()
companyA, _, err := d.Companies().UpsertCompany(fmt.Sprintf("Task Company A %d", suffix), "")
if err != nil {
t.Fatal(err)
}
companyB, _, err := d.Companies().UpsertCompany(fmt.Sprintf("Task Company B %d", suffix), "")
if err != nil {
t.Fatal(err)
}
emptyCompany, _, err := d.Companies().UpsertCompany(fmt.Sprintf("Task Empty Company %d", suffix), "")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
_, _ = d.Exec(`DELETE FROM companies WHERE id IN ($1,$2,$3)`, companyA, companyB, emptyCompany)
})
existingTask, err := d.CreateTask("existing company asset owner", "goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(existingTask.ID) })
var companyAssetA, companyAssetB int64
domainA := fmt.Sprintf("task-company-a-%d.example.test", suffix)
domainB := fmt.Sprintf("task-company-b-%d.example.test", suffix)
if err := d.QueryRow(`
INSERT INTO assets(type, domain, root_domain, company_id, company_source, task_ids)
VALUES ('root_domain',$1,$1,$2,'explicit',ARRAY[$3]::bigint[])
RETURNING id`, domainA, companyA, existingTask.ID).Scan(&companyAssetA); err != nil {
t.Fatal(err)
}
if err := d.QueryRow(`
INSERT INTO assets(type, domain, root_domain, company_id, company_source)
VALUES ('root_domain',$1,$1,$2,'explicit')
RETURNING id`, domainB, companyB).Scan(&companyAssetB); err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _, _ = d.Assets().DeleteByIDs([]int64{companyAssetA, companyAssetB}) })
task, err := d.CreateTaskWithOptions("company-scoped task", "use company scope", TaskCreateOptions{
CompanyIDs: []int64{companyA, companyB},
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(task.ID) })
if fmt.Sprint(task.CompanyIDs) != fmt.Sprint([]int64{companyA, companyB}) {
t.Fatalf("creation result company IDs=%v", task.CompanyIDs)
}
got, err := d.GetTask(task.ID)
if err != nil {
t.Fatal(err)
}
if fmt.Sprint(got.CompanyIDs) != fmt.Sprint([]int64{companyA, companyB}) {
t.Fatalf("hydrated company IDs=%v", got.CompanyIDs)
}
var scopeCount int
if err := d.QueryRow(`SELECT count(*) FROM task_scope WHERE task_id=$1 AND kind='company'`, task.ID).Scan(&scopeCount); err != nil {
t.Fatal(err)
}
if scopeCount != 2 {
t.Fatalf("company task scope rows=%d want 2", scopeCount)
}
assets, err := d.Assets().QueryByTask(task.ID, "root_domain", 10, 0)
if err != nil {
t.Fatal(err)
}
if len(assets) != 2 {
t.Fatalf("company task assets=%d want 2: %+v", len(assets), assets)
}
assetByID := make(map[int64]*Asset, len(assets))
for _, asset := range assets {
assetByID[asset.ID] = asset
}
for assetID, companyName := range map[int64]string{
companyAssetA: fmt.Sprintf("Task Company A %d", suffix),
companyAssetB: fmt.Sprintf("Task Company B %d", suffix),
} {
asset := assetByID[assetID]
if asset == nil {
t.Errorf("company asset %d missing from task", assetID)
continue
}
if asset.TaskSource != taskCompanyAssetSource || asset.TaskSourceSummary != "작업 생성 시 연결된 회사: "+companyName {
t.Errorf("asset %d provenance=%q/%q", assetID, asset.TaskSource, asset.TaskSourceSummary)
}
}
var existingAssociation bool
if err := d.QueryRow(`SELECT $1=ANY(task_ids) FROM assets WHERE id=$2`, existingTask.ID, companyAssetA).Scan(&existingAssociation); err != nil {
t.Fatal(err)
}
if !existingAssociation {
t.Fatal("company snapshot removed the asset's existing task association")
}
var existingSource string
if err := d.QueryRow(`SELECT source FROM task_asset_links WHERE task_id=$1 AND asset_id=$2`, existingTask.ID, companyAssetA).Scan(&existingSource); err != nil {
t.Fatal(err)
}
if existingSource == taskCompanyAssetSource {
t.Fatalf("company snapshot overwrote another task's provenance: %q", existingSource)
}
emptyTask, err := d.CreateTaskWithOptions("empty company scope", "no current assets", TaskCreateOptions{
CompanyIDs: []int64{emptyCompany},
})
if err != nil {
t.Fatal(err)
}
if assets, err := d.Assets().QueryByTask(emptyTask.ID, "", 10, 0); err != nil || len(assets) != 0 {
t.Fatalf("empty company task assets=%+v err=%v", assets, err)
}
if err := d.DeleteTask(emptyTask.ID); err != nil {
t.Fatal(err)
}
duplicateTask, err := d.CreateTaskWithOptions("duplicate company scope", "deduplicate", TaskCreateOptions{
CompanyIDs: []int64{companyA, companyA, companyB, companyA},
})
if err != nil {
t.Fatal(err)
}
if fmt.Sprint(duplicateTask.CompanyIDs) != fmt.Sprint([]int64{companyA, companyB}) {
t.Fatalf("company IDs were not normalized: %v", duplicateTask.CompanyIDs)
}
if err := d.DeleteTask(duplicateTask.ID); err != nil {
t.Fatal(err)
}
badDescription := fmt.Sprintf("invalid-company-%d", suffix)
if _, err := d.CreateTaskWithOptions(badDescription, "rollback", TaskCreateOptions{CompanyIDs: []int64{companyA, 1 << 62}}); !errors.Is(err, ErrTaskCompanyNotFound) {
t.Fatalf("missing company error=%v, want %v", err, ErrTaskCompanyNotFound)
}
var leaked int
if err := d.QueryRow(`SELECT count(*) FROM explorations WHERE description=$1`, badDescription).Scan(&leaked); err != nil {
t.Fatal(err)
}
if leaked != 0 {
t.Fatalf("failed company association leaked %d exploration rows", leaked)
}
}
// TestListTasksOrderByIDDesc pins list ordering to id-descending. created_at is
// deliberately not the sort key: tasks created in the same instant share a
// timestamp and would reorder between polls; id is unique and monotonic.
func TestListTasksOrderByIDDesc(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) — skipping", err)
}
defer d.Close()
stamp := time.Now().UnixNano()
var ids []int64
for i := range 3 {
tk, err := d.CreateTask(fmt.Sprintf("order-%d-%d", stamp, i), "goal", nil, 0, 0)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = d.DeleteTask(tk.ID) })
ids = append(ids, tk.ID)
}
list, err := d.ListTasks()
if err != nil {
t.Fatal(err)
}
// Reduce to just the ids created here; other rows may exist in the shared DB.
mine := map[int64]bool{ids[0]: true, ids[1]: true, ids[2]: true}
var seen []int64
for _, task := range list {
if mine[task.ID] {
seen = append(seen, task.ID)
}
}
if len(seen) != 3 {
t.Fatalf("found %d of the created tasks in the list, want 3", len(seen))
}
// Newest (largest id) first.
if seen[0] != ids[2] || seen[1] != ids[1] || seen[2] != ids[0] {
t.Fatalf("order=%v, want descending %v", seen, []int64{ids[2], ids[1], ids[0]})
}
}
+32
View File
@@ -0,0 +1,32 @@
package db
import (
"database/sql"
"os"
"testing"
_ "github.com/jackc/pgx/v5/stdlib"
)
// TestMain acquires a PostgreSQL advisory lock (7337741002) for the entire
// db test suite. Packages agent and server hold the same lock, so parallel
// `go test ./...` runs serialize across packages on the shared dev DB and
// avoid cross-package cleanup races (e.g. DELETE FROM assets WHERE id > X
// from one package deleting assets created by another).
func TestMain(m *testing.M) {
dsn, _, err := DSN()
if err != nil {
// No DB configured — tests that need PG will skip themselves.
os.Exit(m.Run())
}
conn, err := sql.Open("pgx", dsn)
if err != nil || conn.Ping() != nil {
os.Exit(m.Run())
}
defer conn.Close()
if _, err := conn.Exec(`SELECT pg_advisory_lock(7337741002)`); err != nil {
os.Exit(m.Run())
}
defer conn.Exec(`SELECT pg_advisory_unlock(7337741002)`) //nolint:errcheck
os.Exit(m.Run())
}
+58
View File
@@ -0,0 +1,58 @@
package db
// ToolUsage is one catalog tool invocation. It stores attribution dimensions only:
// tool arguments and results are deliberately excluded from the ledger.
type ToolUsage struct {
ToolKey string `json:"tool_key"`
AgentKey string `json:"agent_key"`
TaskID int64 `json:"task_id"`
ExplorationID int64 `json:"exploration_id"`
IntentID int64 `json:"intent_id"`
SessionID string `json:"session_id"`
}
// InsertToolUsage appends one ledger row. Runtime callers treat metering as
// best-effort so a statistics failure never interrupts the tool itself.
func (d *DB) InsertToolUsage(u *ToolUsage) error {
_, err := d.Exec(`
INSERT INTO tool_usage(tool_key, agent_key, task_id, exploration_id, intent_id, session_id)
VALUES ($1, NULLIF($2,''), $3, $4, $5, NULLIF($6,''))`,
u.ToolKey, u.AgentKey, nullIfZero(u.TaskID), nullIfZero(u.ExplorationID),
nullIfZero(u.IntentID), u.SessionID)
return err
}
// ToolUsageCounts returns invocation totals keyed by catalog tool key. Entries
// without calls are absent; API callers merge the result into the tools catalog.
func (d *DB) ToolUsageCounts() (map[string]int, error) {
rows, err := d.Query(`
SELECT tool_key, COUNT(*)
FROM tool_usage
GROUP BY tool_key`)
if err != nil {
return nil, err
}
defer rows.Close()
out := map[string]int{}
for rows.Next() {
var key string
var calls int
if err := rows.Scan(&key, &calls); err != nil {
return nil, err
}
out[key] = calls
}
if err := rows.Err(); err != nil {
return nil, err
}
archived, err := d.archivedTaskAggregates()
if err != nil {
return nil, err
}
for _, aggregate := range archived {
for key, calls := range aggregate.Tools {
out[key] += calls
}
}
return out, nil
}
+44
View File
@@ -0,0 +1,44 @@
package db
import "testing"
func TestToolUsageLedger(t *testing.T) {
d, err := Open(testDSN(t))
if err != nil {
t.Skipf("postgres unavailable (%v) - skipping", err)
}
defer d.Close()
const (
toolA = "zz_test_tool_usage_a"
toolB = "zz_test_tool_usage_b"
)
cleanup := func() {
_, _ = d.Exec(`DELETE FROM tool_usage WHERE tool_key IN ($1,$2)`, toolA, toolB)
}
cleanup()
defer cleanup()
rows := []*ToolUsage{
{ToolKey: toolA, AgentKey: "worker", TaskID: 991, ExplorationID: 5, IntentID: 7},
{ToolKey: toolA, AgentKey: "planner", TaskID: 992, ExplorationID: 6},
{ToolKey: toolA, AgentKey: "chatbot", SessionID: "conv-1"},
{ToolKey: toolB, AgentKey: "worker", TaskID: 991},
}
for _, row := range rows {
if err := d.InsertToolUsage(row); err != nil {
t.Fatalf("insert %s: %v", row.ToolKey, err)
}
}
counts, err := d.ToolUsageCounts()
if err != nil {
t.Fatal(err)
}
if got := counts[toolA]; got != 3 {
t.Errorf("%s calls: want 3, got %d", toolA, got)
}
if got := counts[toolB]; got != 1 {
t.Errorf("%s calls: want 1, got %d", toolB, got)
}
}
+232
View File
@@ -0,0 +1,232 @@
package db
import (
"database/sql"
"encoding/json"
)
// Tool is one row of the built-in tool catalog. key + handler live in code; this
// row carries only the page-editable surface: description, parameter schema
// (structure read-only, per-param description/default editable), agent binding,
// and the on/off switch. See schema.sql §H and agent/toolcatalog.go.
type Tool struct {
Key string `json:"key"`
System bool `json:"system"`
Description string `json:"description"`
Schema json.RawMessage `json:"schema"`
Agents []string `json:"agents"`
Enabled bool `json:"enabled"`
Kind string `json:"kind"` // builtin | command | script | http
Exec json.RawMessage `json:"exec"` // 自定义工具执行规格(kind!=builtin)
Deferred bool `json:"deferred"` // schema 延迟(走 SearchExtraTools/ExecuteExtraTool)
Calls int `json:"calls"` // runtime ledger aggregate; not stored in tools
}
const toolCols = `key, system, description, schema, agents, enabled, kind, exec, deferred`
// SeedTool inserts a built-in tool's code-defined defaults ONCE. ON CONFLICT DO
// NOTHING: an existing row (possibly edited in the UI) is never overwritten on
// startup — that's what keeps page edits from being wiped every restart. Use
// UpsertToolForce for an explicit "reset to code default".
func (d *DB) SeedTool(key, desc string, schema, agents json.RawMessage) error {
if len(schema) == 0 {
schema = json.RawMessage("{}")
}
if len(agents) == 0 {
agents = json.RawMessage("[]")
}
_, err := d.Exec(`
INSERT INTO tools(key, system, description, schema, agents, enabled)
VALUES ($1, true, $2, $3, $4, true)
ON CONFLICT (key) DO NOTHING`, key, desc, schema, agents)
return err
}
// AddAgentToToolBinding adds an agent key to the given tools' `agents` arrays if
// not already present (idempotent). Used to give the built-in Auto agent its
// default toolset on existing DBs.
func (d *DB) AddAgentToToolBinding(agentKey string, keys []string) error {
for _, k := range keys {
if _, err := d.Exec(`UPDATE tools SET agents = agents || to_jsonb($1::text) WHERE key=$2 AND NOT (agents ? $1)`, agentKey, k); err != nil {
return err
}
}
return nil
}
// RemoveAgentFromToolBindings strips an agent key from every tool's `agents`
// JSONB array — called when a custom agent is deleted so no tool keeps a dangling
// binding. Uses jsonb `-` (remove array element) guarded by `?` (membership).
func (d *DB) RemoveAgentFromToolBindings(agentKey string) error {
_, err := d.Exec(`UPDATE tools SET agents = agents - $1 WHERE agents ? $1`, agentKey)
return err
}
// RemoveAgentFromTool strips one agent key from a SINGLE tool's `agents` array —
// used by one-time migrations that change a tool's default binding on existing DBs
// (SeedTool is first-insert-only, so a changed default never reaches a seeded row).
func (d *DB) RemoveAgentFromTool(agentKey, toolKey string) error {
_, err := d.Exec(`UPDATE tools SET agents = agents - $1 WHERE key=$2 AND agents ? $1`, agentKey, toolKey)
return err
}
// UpsertToolForce overwrites a tool row with the given code-default values (used by
// the per-tool "reset" action). It resets description/schema/agents and re-enables
// the tool, but keeps system=true.
func (d *DB) UpsertToolForce(key, desc string, schema, agents json.RawMessage) error {
if len(schema) == 0 {
schema = json.RawMessage("{}")
}
if len(agents) == 0 {
agents = json.RawMessage("[]")
}
_, err := d.Exec(`
INSERT INTO tools(key, system, description, schema, agents, enabled)
VALUES ($1, true, $2, $3, $4, true)
ON CONFLICT (key) DO UPDATE
SET description = EXCLUDED.description,
schema = EXCLUDED.schema,
agents = EXCLUDED.agents,
enabled = true`, key, desc, schema, agents)
return err
}
// RefreshToolDefaults updates a system tool's model-facing description + schema to the
// code defaults, PRESERVING the user's agent binding + enabled flag. Used by one-time
// migrations to propagate a code schema change (SeedTool is first-insert-only, so a new
// parameter added in code otherwise never reaches an already-seeded row). No-op for
// custom tools or unknown keys.
func (d *DB) RefreshToolDefaults(key, desc string, schema json.RawMessage) error {
if len(schema) == 0 {
schema = json.RawMessage("{}")
}
_, err := d.Exec(`UPDATE tools SET description=$2, schema=$3, updated_at=now() WHERE key=$1 AND system`, key, desc, schema)
return err
}
func scanTool(rows interface{ Scan(...any) error }) (*Tool, error) {
var t Tool
var agents []byte
if err := rows.Scan(&t.Key, &t.System, &t.Description, &t.Schema, &agents, &t.Enabled, &t.Kind, &t.Exec, &t.Deferred); err != nil {
return nil, err
}
if len(agents) > 0 {
_ = json.Unmarshal(agents, &t.Agents)
}
if t.Agents == nil {
t.Agents = []string{}
}
return &t, nil
}
// ListTools returns the whole tool catalog, ordered by key.
func (d *DB) ListTools() ([]*Tool, error) {
rows, err := d.Query(`SELECT ` + toolCols + ` FROM tools ORDER BY key`)
if err != nil {
return nil, err
}
defer rows.Close()
var out []*Tool
for rows.Next() {
t, err := scanTool(rows)
if err != nil {
return nil, err
}
out = append(out, t)
}
return out, rows.Err()
}
// GetTool fetches one tool row (nil if absent).
func (d *DB) GetTool(key string) (*Tool, error) {
row := d.QueryRow(`SELECT `+toolCols+` FROM tools WHERE key=$1`, key)
t, err := scanTool(row)
if err == sql.ErrNoRows {
return nil, nil
}
return t, err
}
// CreateCustomTool inserts a user-defined tool (system=false) with an execution
// spec. Fails if the key already exists.
func (d *DB) CreateCustomTool(t *Tool) error {
schema := t.Schema
if len(schema) == 0 {
schema = json.RawMessage("{}")
}
exec := t.Exec
if len(exec) == 0 {
exec = json.RawMessage("{}")
}
agents, _ := json.Marshal(t.Agents)
if len(agents) == 0 {
agents = json.RawMessage("[]")
}
_, err := d.Exec(`
INSERT INTO tools(key, system, description, schema, agents, enabled, kind, exec, deferred)
VALUES ($1, false, $2, $3, $4, $5, $6, $7, $8)`,
t.Key, t.Description, schema, agents, t.Enabled, t.Kind, exec, t.Deferred)
return err
}
// UpdateCustomTool updates a custom tool's editable fields (kind/exec/deferred +
// desc/schema/agents/enabled). Only touches system=false rows.
func (d *DB) UpdateCustomTool(t *Tool) error {
schema := t.Schema
if len(schema) == 0 {
schema = json.RawMessage("{}")
}
exec := t.Exec
if len(exec) == 0 {
exec = json.RawMessage("{}")
}
agents, _ := json.Marshal(t.Agents)
if len(agents) == 0 {
agents = json.RawMessage("[]")
}
_, err := d.Exec(`
UPDATE tools SET description=$2, schema=$3, agents=$4, enabled=$5, kind=$6, exec=$7, deferred=$8
WHERE key=$1 AND system=false`,
t.Key, t.Description, schema, agents, t.Enabled, t.Kind, exec, t.Deferred)
return err
}
// DeleteCustomTool removes a custom tool (system=false only; built-ins protected).
func (d *DB) DeleteCustomTool(key string) error {
_, err := d.Exec(`DELETE FROM tools WHERE key=$1 AND system=false`, key)
return err
}
// ListCustomTools returns only the user-defined (system=false) tools.
func (d *DB) ListCustomTools() ([]*Tool, error) {
rows, err := d.Query(`SELECT ` + toolCols + ` FROM tools WHERE system=false ORDER BY key`)
if err != nil {
return nil, err
}
defer rows.Close()
var out []*Tool
for rows.Next() {
t, err := scanTool(rows)
if err != nil {
return nil, err
}
out = append(out, t)
}
return out, rows.Err()
}
// UpdateTool saves the page-editable fields. key is never changed (it is welded to
// the Go handler). system tools: the caller must keep the schema structure — only
// per-param description/default and the agent binding are meant to move.
func (d *DB) UpdateTool(key, desc string, schema, agents json.RawMessage, enabled bool) error {
if len(schema) == 0 {
schema = json.RawMessage("{}")
}
if len(agents) == 0 {
agents = json.RawMessage("[]")
}
_, err := d.Exec(`
UPDATE tools SET description=$2, schema=$3, agents=$4, enabled=$5 WHERE key=$1`,
key, desc, schema, agents, enabled)
return err
}
+299
View File
@@ -0,0 +1,299 @@
package db
import (
"database/sql"
"encoding/json"
"time"
)
// AgentTrigger is one P3 trigger attached to a custom agent. Six trigger
// conditions can be on at once: interval / on_finding / on_goal_met /
// on_task_timeout / on_tool_call / on_task_create. ToolNames scopes the tool-call
// trigger to a non-empty set of tool keys (empty is rejected at the API layer for
// on_tool_call).
type AgentTrigger struct {
ID int64 `json:"id"`
AgentKey string `json:"agent_key"`
Enabled bool `json:"enabled"`
IntervalSec int `json:"interval_sec"`
OnFinding bool `json:"on_finding"`
OnGoalMet bool `json:"on_goal_met"`
OnTaskTimeout bool `json:"on_task_timeout"`
OnToolCall bool `json:"on_tool_call"`
OnTaskCreate bool `json:"on_task_create"`
IntervalMessage string `json:"interval_message"` // 各触发条件的独立用户消息
FindingMessage string `json:"finding_message"`
GoalMessage string `json:"goal_message"`
TaskTimeoutMessage string `json:"task_timeout_message"`
ToolCallMessage string `json:"tool_call_message"`
TaskCreateMessage string `json:"task_create_message"`
ToolNames []string `json:"tool_names"` // 选中的工具 key 列表(DB 存 JSON 文本)
LastFire *time.Time `json:"last_fire,omitempty"`
}
const triggerCols = `id, agent_key, enabled, interval_sec, on_finding, on_goal_met, on_task_timeout, on_tool_call, on_task_create, interval_message, finding_message, goal_message, task_timeout_message, tool_call_message, task_create_message, tool_names, last_fire`
// marshalToolNames encodes the tool-key list as JSON text for the tool_names column.
// A nil/empty list stores "" (not "null"/"[]") so the column default stays clean.
func marshalToolNames(names []string) string {
if len(names) == 0 {
return ""
}
b, err := json.Marshal(names)
if err != nil {
return ""
}
return string(b)
}
func scanTrigger(sc interface{ Scan(...any) error }) (*AgentTrigger, error) {
var t AgentTrigger
var lf sql.NullTime
var toolNames string
if err := sc.Scan(&t.ID, &t.AgentKey, &t.Enabled, &t.IntervalSec, &t.OnFinding, &t.OnGoalMet, &t.OnTaskTimeout, &t.OnToolCall, &t.OnTaskCreate,
&t.IntervalMessage, &t.FindingMessage, &t.GoalMessage, &t.TaskTimeoutMessage, &t.ToolCallMessage, &t.TaskCreateMessage, &toolNames, &lf); err != nil {
return nil, err
}
t.ToolNames = []string{}
if toolNames != "" {
_ = json.Unmarshal([]byte(toolNames), &t.ToolNames)
}
if lf.Valid {
t.LastFire = &lf.Time
}
return &t, nil
}
// CreateTrigger inserts a trigger for agentKey and returns it.
func (d *DB) CreateTrigger(t *AgentTrigger) (*AgentTrigger, error) {
row := d.QueryRow(`
INSERT INTO agent_triggers(agent_key, enabled, interval_sec, on_finding, on_goal_met, on_task_timeout, on_tool_call, on_task_create, interval_message, finding_message, goal_message, task_timeout_message, tool_call_message, task_create_message, tool_names)
VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15) RETURNING `+triggerCols,
t.AgentKey, t.Enabled, t.IntervalSec, t.OnFinding, t.OnGoalMet, t.OnTaskTimeout, t.OnToolCall, t.OnTaskCreate,
t.IntervalMessage, t.FindingMessage, t.GoalMessage, t.TaskTimeoutMessage, t.ToolCallMessage, t.TaskCreateMessage, marshalToolNames(t.ToolNames))
return scanTrigger(row)
}
// UpdateTrigger updates a trigger's fields (not last_fire).
func (d *DB) UpdateTrigger(t *AgentTrigger) error {
_, err := d.Exec(`UPDATE agent_triggers SET enabled=$2, interval_sec=$3, on_finding=$4, on_goal_met=$5, on_task_timeout=$6, on_tool_call=$7, on_task_create=$8, interval_message=$9, finding_message=$10, goal_message=$11, task_timeout_message=$12, tool_call_message=$13, task_create_message=$14, tool_names=$15 WHERE id=$1`,
t.ID, t.Enabled, t.IntervalSec, t.OnFinding, t.OnGoalMet, t.OnTaskTimeout, t.OnToolCall, t.OnTaskCreate,
t.IntervalMessage, t.FindingMessage, t.GoalMessage, t.TaskTimeoutMessage, t.ToolCallMessage, t.TaskCreateMessage, marshalToolNames(t.ToolNames))
return err
}
// DeleteTrigger removes a trigger.
func (d *DB) DeleteTrigger(id int64) error {
_, err := d.Exec(`DELETE FROM agent_triggers WHERE id=$1`, id)
return err
}
// ListTriggersFor returns an agent's triggers.
func (d *DB) ListTriggersFor(agentKey string) ([]*AgentTrigger, error) {
return d.queryTriggers(`SELECT `+triggerCols+` FROM agent_triggers WHERE agent_key=$1 ORDER BY id`, agentKey)
}
// ListEnabledTriggers returns all enabled triggers (for the scheduler).
func (d *DB) ListEnabledTriggers() ([]*AgentTrigger, error) {
return d.queryTriggers(`SELECT ` + triggerCols + ` FROM agent_triggers WHERE enabled ORDER BY id`)
}
func (d *DB) queryTriggers(q string, args ...any) ([]*AgentTrigger, error) {
rows, err := d.Query(q, args...)
if err != nil {
return nil, err
}
defer rows.Close()
out := []*AgentTrigger{}
for rows.Next() {
t, err := scanTrigger(rows)
if err != nil {
return nil, err
}
out = append(out, t)
}
return out, rows.Err()
}
// TouchTriggerFire records an interval trigger's fire time (now).
func (d *DB) TouchTriggerFire(id int64) error {
_, err := d.Exec(`UPDATE agent_triggers SET last_fire=now() WHERE id=$1`, id)
return err
}
// DeleteTriggersForAgent removes all triggers of an agent (custom agent delete).
func (d *DB) DeleteTriggersForAgent(agentKey string) error {
_, err := d.Exec(`DELETE FROM agent_triggers WHERE agent_key=$1`, agentKey)
return err
}
// ---------- scheduler_state (kv watermarks) ----------
func (d *DB) GetSchedState(key string) (string, error) {
var v string
err := d.QueryRow(`SELECT value FROM scheduler_state WHERE key=$1`, key).Scan(&v)
if err == sql.ErrNoRows {
return "", nil
}
return v, err
}
func (d *DB) SetSchedState(key, value string) error {
_, err := d.Exec(`INSERT INTO scheduler_state(key,value) VALUES ($1,$2)
ON CONFLICT (key) DO UPDATE SET value=EXCLUDED.value`, key, value)
return err
}
// ---------- event queries (cross-exploration, for the scheduler) ----------
// TaskEvent is a finding/goal event carrying the owning task's info, used to
// compose the trigger message context.
type TaskEvent struct {
NodeID int64 `json:"node_id"`
TaskID int64 `json:"task_id"`
TaskDesc string `json:"task_description"`
TaskGoal string `json:"task_goal"`
Summary string `json:"summary"` // finding summary / goal text
VulnClass string `json:"vulnclass"` // finding only
Severity string `json:"severity"` // finding only
Tool string `json:"tool"` // tool-call only: tool name
ToolInput string `json:"tool_input"` // tool-call only: 入参(JSON 文本)
ToolOutput string `json:"tool_output"` // tool-call only: 返回内容
ToolIsErr bool `json:"tool_is_err"` // tool-call only: 工具返回是否为错误
}
// NewFindingsSince returns findings with node id > lastID across all live tasks,
// ordered by id (monotonic watermark → no double-fire).
func (d *DB) NewFindingsSince(lastID int64) ([]TaskEvent, error) {
rows, err := d.Query(`
SELECT n.id, t.id, t.description, t.goal, n.payload
FROM exploration_nodes n JOIN tasks t ON t.exploration_id = n.exploration_id
WHERE n.kind='finding' AND n.id > $1 AND t.deleted_at IS NULL
ORDER BY n.id`, lastID)
if err != nil {
return nil, err
}
defer rows.Close()
out := []TaskEvent{}
for rows.Next() {
var e TaskEvent
var payload []byte
if err := rows.Scan(&e.NodeID, &e.TaskID, &e.TaskDesc, &e.TaskGoal, &payload); err != nil {
return nil, err
}
var p struct{ Summary, Vulnclass, Severity string }
_ = json.Unmarshal(payload, &p)
e.Summary, e.VulnClass, e.Severity = p.Summary, p.Vulnclass, p.Severity
out = append(out, e)
}
return out, rows.Err()
}
// TimedOutTasksSince returns tasks that reached status='timeout' with id > lastID,
// ordered by id (monotonic watermark → no double-fire across restarts).
func (d *DB) TimedOutTasksSince(lastID int64) ([]TaskEvent, error) {
rows, err := d.Query(`
SELECT id, description, goal FROM tasks
WHERE status='timeout' AND deleted_at IS NULL AND id > $1
ORDER BY id`, lastID)
if err != nil {
return nil, err
}
defer rows.Close()
out := []TaskEvent{}
for rows.Next() {
var e TaskEvent
if err := rows.Scan(&e.NodeID, &e.TaskDesc, &e.TaskGoal); err != nil {
return nil, err
}
e.TaskID = e.NodeID // task id doubles as NodeID for the watermark
e.Summary = e.TaskGoal
out = append(out, e)
}
return out, rows.Err()
}
// NewTasksSince returns tasks created with id > lastID (excluding deleted),
// ordered by id (monotonic watermark → no double-fire across restarts). Triggered
// agent runs are conversations, not tasks, so this never fires on its own output.
func (d *DB) NewTasksSince(lastID int64) ([]TaskEvent, error) {
rows, err := d.Query(`
SELECT id, description, goal FROM tasks
WHERE deleted_at IS NULL AND id > $1
ORDER BY id`, lastID)
if err != nil {
return nil, err
}
defer rows.Close()
out := []TaskEvent{}
for rows.Next() {
var e TaskEvent
if err := rows.Scan(&e.NodeID, &e.TaskDesc, &e.TaskGoal); err != nil {
return nil, err
}
e.TaskID = e.NodeID // task id doubles as NodeID for the watermark
e.Summary = e.TaskGoal
out = append(out, e)
}
return out, rows.Err()
}
// NewToolCallsSince returns completed tool calls (a tool_result row) with activity
// id > lastID across all live tasks, ordered by id (monotonic watermark → no
// double-fire). It is driven by tool_result rows (the tool finished, so both input
// and output are available) and joins back to the paired tool_use row for the input.
// Only task-execution activity is scanned — triggered agent runs are conversations
// (conversation_activities), so a tool-call trigger never fires on its own output.
func (d *DB) NewToolCallsSince(lastID int64) ([]TaskEvent, error) {
rows, err := d.Query(`
SELECT r.id, t.id, t.description, t.goal, r.tool, COALESCE(u.detail,''), COALESCE(r.detail,''), r.is_error
FROM activity r
JOIN tasks t ON t.exploration_id = r.exploration_id
LEFT JOIN activity u ON u.exploration_id = r.exploration_id AND u.tool_use_id = r.tool_use_id AND u.kind='tool_use'
WHERE r.kind='tool_result' AND r.id > $1 AND r.tool <> '' AND t.deleted_at IS NULL
ORDER BY r.id`, lastID)
if err != nil {
return nil, err
}
defer rows.Close()
out := []TaskEvent{}
for rows.Next() {
var e TaskEvent
if err := rows.Scan(&e.NodeID, &e.TaskID, &e.TaskDesc, &e.TaskGoal, &e.Tool, &e.ToolInput, &e.ToolOutput, &e.ToolIsErr); err != nil {
return nil, err
}
out = append(out, e)
}
return out, rows.Err()
}
// MetGoals returns all met goals across live tasks (the scheduler filters out the
// ones it already fired for via the persisted fired-set).
func (d *DB) MetGoals() ([]TaskEvent, error) {
rows, err := d.Query(`
SELECT n.id, t.id, t.description, t.goal, n.payload
FROM exploration_nodes n JOIN tasks t ON t.exploration_id = n.exploration_id
WHERE n.kind='goal' AND n.state='met' AND t.deleted_at IS NULL
ORDER BY n.id`)
if err != nil {
return nil, err
}
defer rows.Close()
out := []TaskEvent{}
for rows.Next() {
var e TaskEvent
var payload []byte
if err := rows.Scan(&e.NodeID, &e.TaskID, &e.TaskDesc, &e.TaskGoal, &payload); err != nil {
return nil, err
}
var p struct{ Text, Summary string }
_ = json.Unmarshal(payload, &p)
if p.Text != "" {
e.Summary = p.Text
} else {
e.Summary = p.Summary
}
out = append(out, e)
}
return out, rows.Err()
}