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:
+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
|
||||
}
|
||||
Reference in New Issue
Block a user