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
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:
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
+1126
File diff suppressed because it is too large
Load Diff
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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(¤t); 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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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()
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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))
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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])
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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 = ©
|
||||
}
|
||||
}
|
||||
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, ¤tRevision); 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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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()
|
||||
}
|
||||
@@ -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]})
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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()
|
||||
}
|
||||
Reference in New Issue
Block a user