package store import ( "fmt" "atlas9.dev/c/demo/lib/filter" ) // sqlField maps a filter field name to its SQL column expression and any // JOIN clause required to reference it. type sqlField struct { Name string Column string // SQL column expression; defaults to Name if empty Join string // JOIN clause added when this field appears in a filter } func (f sqlField) column() string { if f.Column != "" { return f.Column } return f.Name } // filterToSQL compiles a filter expression into a SQL WHERE fragment and the // JOIN clauses needed by the referenced fields. // // Returns ("", nil, args) unchanged when expr is nil. // The WHERE fragment starts with "AND " and uses positional parameters // starting at len(args)+1. func filterToSQL(expr filter.Expr, fields []sqlField, args []any) (where string, joins []string, newArgs []any) { if expr == nil { return "", nil, args } m := make(map[string]sqlField, len(fields)) for _, f := range fields { m[f.Name] = f } b := &filterSQLBuilder{fields: m, args: args} clause := b.build(expr) return "AND " + clause, b.joins, b.args } type filterSQLBuilder struct { fields map[string]sqlField args []any joins []string seen map[string]bool } func (b *filterSQLBuilder) addJoin(join string) { if join == "" { return } if b.seen == nil { b.seen = map[string]bool{} } if !b.seen[join] { b.seen[join] = true b.joins = append(b.joins, join) } } func (b *filterSQLBuilder) arg(v any) string { b.args = append(b.args, v) return fmt.Sprintf("$%d", len(b.args)) } func (b *filterSQLBuilder) build(expr filter.Expr) string { switch e := expr.(type) { case filter.AndExpr: return fmt.Sprintf("(%s AND %s)", b.build(e.Left), b.build(e.Right)) case filter.OrExpr: return fmt.Sprintf("(%s OR %s)", b.build(e.Left), b.build(e.Right)) case filter.NotExpr: return fmt.Sprintf("(NOT %s)", b.build(e.Operand)) case filter.CompareExpr: f := b.fields[e.Field] b.addJoin(f.Join) col := f.column() switch e.Op { case filter.OpEq: return fmt.Sprintf("%s = %s", col, b.arg(e.Value)) case filter.OpNe: return fmt.Sprintf("%s != %s", col, b.arg(e.Value)) case filter.OpLt: return fmt.Sprintf("%s < %s", col, b.arg(e.Value)) case filter.OpGt: return fmt.Sprintf("%s > %s", col, b.arg(e.Value)) case filter.OpLe: return fmt.Sprintf("%s <= %s", col, b.arg(e.Value)) case filter.OpGe: return fmt.Sprintf("%s >= %s", col, b.arg(e.Value)) case filter.OpHas: return fmt.Sprintf("%s LIKE %s", col, b.arg("%"+e.Value+"%")) } } return "TRUE" }