feat(filter): standard CEL now variable, time accessors, set ops
Replace the custom now() function with an idiomatic `now` timestamp variable (host-injected, frozen once per compile) and retype created_ts/updated_ts/create_time to CEL timestamp. Filters now use standard timestamp/duration arithmetic, e.g. `created_ts >= now - duration("24h")` and `timestamp("2025-01-01T00:00:00Z")`.
Add standard CEL surface that compiles to SQL across SQLite/MySQL/Postgres: timestamp accessors (getFullYear/getMonth/getDate/getDayOfWeek/..., with 0-based month and weekday normalized), ext.Sets() (sets.contains/intersects/equivalent over tags), tags.exists_one(), size() on string fields, and division/modulo folding. A frozen clock is injectable for deterministic tests.
BREAKING CHANGE: now() is removed (use the `now` variable) and time fields are timestamps, so bare-epoch comparisons need timestamp(<epoch>). Existing saved shortcuts using the old syntax must be updated.
This commit is contained in:
parent
cafa56f1a8
commit
26f4b73cb9
14 changed files with 943 additions and 122 deletions
|
|
@ -48,6 +48,12 @@ stmt, _ := engine.CompileToStatement(ctx, `has_task_list && visibility == "PUBLI
|
|||
tracks offsets to compose queries with pre-existing arguments.
|
||||
- **JSON Fields** — Memo metadata lives in `memo.payload`. The renderer handles
|
||||
`JSON_EXTRACT`/`json_extract`/`->`/`->>` variations and boolean coercion.
|
||||
- **Time Fields** — `created_ts`, `updated_ts`, and attachment `create_time` are
|
||||
CEL `timestamp` values. Express instants with the `now` variable,
|
||||
`duration("…")` (e.g. `created_ts >= now - duration("24h")`), or
|
||||
`timestamp("2006-01-02T15:04:05Z")` / `timestamp(<epoch-seconds>)`. These fold
|
||||
to epoch seconds at compile time — `now` is frozen once per compile (injectable
|
||||
for tests via the engine clock) — so the backing columns stay unchanged.
|
||||
- **Tag Operations** — `tag in [...]` and `"tag" in tags` become JSON array
|
||||
predicates. SQLite uses `LIKE` patterns, MySQL uses `JSON_CONTAINS`, and
|
||||
Postgres uses `@>`.
|
||||
|
|
@ -64,9 +70,26 @@ stmt, _ := engine.CompileToStatement(ctx, `has_task_list && visibility == "PUBLI
|
|||
Go's RE2 via `cel.ValidateRegexLiterals()`. **Caveat:** regex *syntax* differs
|
||||
per engine (Go RE2 on SQLite, POSIX ERE on Postgres, ICU on MySQL 8.0+), so
|
||||
engine-specific patterns may not be portable.
|
||||
- **Tag `all()`** — `tags.all(t, <pred>)` matches only non-empty tag sets where
|
||||
every element satisfies the predicate, via per-element iteration
|
||||
(`json_each` / `jsonb_array_elements_text` / `JSON_TABLE`).
|
||||
- **Tag `all()` / `exists_one()`** — `tags.all(t, <pred>)` matches only non-empty
|
||||
tag sets where every element satisfies the predicate; `tags.exists_one(t,
|
||||
<pred>)` matches when exactly one element does (`COUNT(...) = 1`). Both iterate
|
||||
per-element (`json_each` / `jsonb_array_elements_text` / `JSON_TABLE`).
|
||||
- **Timestamp Accessors** — `created_ts.getFullYear()`, `getMonth()`, `getDate()`,
|
||||
`getDayOfMonth()`, `getDayOfWeek()`, `getDayOfYear()`, `getHours()`,
|
||||
`getMinutes()`, `getSeconds()` render to date-part extraction (`strftime` /
|
||||
`EXTRACT` / `YEAR`/`MONTH`/…). Results are normalized to CEL's base (0-based
|
||||
month, 0-based day-of-week with 0 = Sunday). Extraction is UTC on SQLite/Postgres
|
||||
(epoch columns); on MySQL the `TIMESTAMP` column is read in the session time
|
||||
zone. A timezone argument is not supported.
|
||||
- **Set Operations** — `ext.Sets()`: `sets.contains(tags, [...])`,
|
||||
`sets.intersects(tags, [...])`, and `sets.equivalent(tags, [...])` desugar to
|
||||
exact-membership checks (AND / OR of `"v" in tags`); `equivalent` adds a
|
||||
`size(tags)` length check (relies on tags being a set).
|
||||
- **`size()`** — `size(tags)` renders to JSON array length; `size(content)` (and
|
||||
other string fields) render to `LENGTH` / `CHAR_LENGTH` (MySQL) for code-point
|
||||
counts.
|
||||
- **Arithmetic** — `+`, `-`, `*`, `/`, `%` constant-fold on literal/`now`/`duration`
|
||||
operands (division and modulo guard against a zero divisor).
|
||||
|
||||
## Typical Integration
|
||||
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import (
|
|||
"context"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/google/cel-go/cel"
|
||||
"github.com/pkg/errors"
|
||||
|
|
@ -13,6 +14,10 @@ import (
|
|||
type Engine struct {
|
||||
schema Schema
|
||||
env *cel.Env
|
||||
// nowFunc resolves the value of the `now` variable. It is frozen once per
|
||||
// Compile so a single filter sees a single instant, and is overridable in
|
||||
// tests for deterministic folding.
|
||||
nowFunc func() time.Time
|
||||
}
|
||||
|
||||
// NewEngine builds a new Engine for the provided schema.
|
||||
|
|
@ -22,8 +27,9 @@ func NewEngine(schema Schema) (*Engine, error) {
|
|||
return nil, errors.Wrap(err, "failed to create CEL environment")
|
||||
}
|
||||
return &Engine{
|
||||
schema: schema,
|
||||
env: env,
|
||||
schema: schema,
|
||||
env: env,
|
||||
nowFunc: time.Now,
|
||||
}, nil
|
||||
}
|
||||
|
||||
|
|
@ -53,7 +59,7 @@ func (e *Engine) Compile(_ context.Context, filter string) (*Program, error) {
|
|||
return nil, errors.Wrap(err, "failed to convert AST")
|
||||
}
|
||||
|
||||
cond, err := buildCondition(parsed.GetExpr(), e.schema)
|
||||
cond, err := buildCondition(parsed.GetExpr(), parseContext{schema: e.schema, now: e.nowFunc()})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
|
|||
192
internal/filter/functions_test.go
Normal file
192
internal/filter/functions_test.go
Normal file
|
|
@ -0,0 +1,192 @@
|
|||
package filter
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Arithmetic folding: division and modulo
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCompileDivisionFolds(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine, err := NewEngine(NewSchema())
|
||||
require.NoError(t, err)
|
||||
stmt, err := engine.CompileToStatement(context.Background(), `creator_id == 100 / 10`, RenderOptions{Dialect: DialectSQLite})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []any{int64(10)}, stmt.Args)
|
||||
}
|
||||
|
||||
func TestCompileModuloFolds(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine, err := NewEngine(NewSchema())
|
||||
require.NoError(t, err)
|
||||
stmt, err := engine.CompileToStatement(context.Background(), `creator_id == 17 % 5`, RenderOptions{Dialect: DialectSQLite})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []any{int64(2)}, stmt.Args)
|
||||
}
|
||||
|
||||
func TestCompileDivisionByZeroErrors(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine, err := NewEngine(NewSchema())
|
||||
require.NoError(t, err)
|
||||
_, err = engine.Compile(context.Background(), `creator_id == 10 / 0`)
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// size() on scalar string fields -> SQL length
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCompileSizeOnContentRendersLengthPerDialect(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine, err := NewEngine(NewSchema())
|
||||
require.NoError(t, err)
|
||||
|
||||
cases := []struct {
|
||||
dialect DialectName
|
||||
fragment string
|
||||
}{
|
||||
{DialectSQLite, "LENGTH("},
|
||||
{DialectPostgres, "LENGTH("},
|
||||
{DialectMySQL, "CHAR_LENGTH("},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
stmt, err := engine.CompileToStatement(context.Background(), `size(content) > 5`, RenderOptions{Dialect: tc.dialect})
|
||||
require.NoError(t, err, tc.dialect)
|
||||
require.Contains(t, stmt.SQL, tc.fragment, "dialect %s", tc.dialect)
|
||||
require.Equal(t, []any{int64(5)}, stmt.Args, "dialect %s", tc.dialect)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Timestamp accessor methods (getFullYear, getMonth, ...)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCompileTimestampAccessorsPerDialect(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine, err := NewEngine(NewSchema())
|
||||
require.NoError(t, err)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
filter string
|
||||
dialect DialectName
|
||||
fragments []string
|
||||
arg int64
|
||||
}{
|
||||
// getFullYear == 2024
|
||||
{"sqlite year", `created_ts.getFullYear() == 2024`, DialectSQLite, []string{"strftime('%Y'", "'unixepoch'"}, 2024},
|
||||
{"pg year", `created_ts.getFullYear() == 2024`, DialectPostgres, []string{"EXTRACT(YEAR FROM to_timestamp(", "AT TIME ZONE 'UTC'"}, 2024},
|
||||
{"mysql year", `created_ts.getFullYear() == 2024`, DialectMySQL, []string{"YEAR(`memo`.`created_ts`)"}, 2024},
|
||||
// getMonth is 0-based -> SQL must subtract 1
|
||||
{"sqlite month", `created_ts.getMonth() == 5`, DialectSQLite, []string{"strftime('%m'", "- 1)"}, 5},
|
||||
{"pg month", `created_ts.getMonth() == 5`, DialectPostgres, []string{"EXTRACT(MONTH FROM", "- 1)"}, 5},
|
||||
{"mysql month", `created_ts.getMonth() == 5`, DialectMySQL, []string{"MONTH(`memo`.`created_ts`)", "- 1)"}, 5},
|
||||
// getDayOfWeek 0=Sunday -> MySQL DAYOFWEEK is 1-based and must subtract 1
|
||||
{"mysql dow", `created_ts.getDayOfWeek() == 0`, DialectMySQL, []string{"DAYOFWEEK(`memo`.`created_ts`)", "- 1)"}, 0},
|
||||
{"sqlite dow", `created_ts.getDayOfWeek() == 0`, DialectSQLite, []string{"strftime('%w'"}, 0},
|
||||
// getDate is 1-based -> no offset
|
||||
{"sqlite date", `created_ts.getDate() == 22`, DialectSQLite, []string{"strftime('%d'"}, 22},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
stmt, err := engine.CompileToStatement(context.Background(), tc.filter, RenderOptions{Dialect: tc.dialect})
|
||||
require.NoError(t, err, tc.name)
|
||||
for _, frag := range tc.fragments {
|
||||
require.Contains(t, stmt.SQL, frag, tc.name)
|
||||
}
|
||||
require.Equal(t, []any{tc.arg}, stmt.Args, tc.name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompileTimestampAccessorRejectsTimezoneArg(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine, err := NewEngine(NewSchema())
|
||||
require.NoError(t, err)
|
||||
_, err = engine.Compile(context.Background(), `created_ts.getHours("America/New_York") == 9`)
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
func TestCompileTimestampAccessorRejectsNonTimestampField(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine, err := NewEngine(NewSchema())
|
||||
require.NoError(t, err)
|
||||
// content is a string, not a timestamp.
|
||||
_, err = engine.Compile(context.Background(), `content.getFullYear() == 2024`)
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// ext.Sets(): sets.contains / sets.intersects / sets.equivalent over tags
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCompileSetsContainsRendersAndOfMemberships(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine, err := NewEngine(NewSchema())
|
||||
require.NoError(t, err)
|
||||
stmt, err := engine.CompileToStatement(context.Background(), `sets.contains(tags, ["a", "b"])`, RenderOptions{Dialect: DialectSQLite})
|
||||
require.NoError(t, err)
|
||||
require.Contains(t, stmt.SQL, " AND ")
|
||||
require.Len(t, stmt.Args, 2)
|
||||
}
|
||||
|
||||
func TestCompileSetsIntersectsRendersOrOfMemberships(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine, err := NewEngine(NewSchema())
|
||||
require.NoError(t, err)
|
||||
stmt, err := engine.CompileToStatement(context.Background(), `sets.intersects(tags, ["a", "b"])`, RenderOptions{Dialect: DialectSQLite})
|
||||
require.NoError(t, err)
|
||||
require.Contains(t, stmt.SQL, " OR ")
|
||||
require.Len(t, stmt.Args, 2)
|
||||
}
|
||||
|
||||
func TestCompileSetsEquivalentAddsLengthCheck(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine, err := NewEngine(NewSchema())
|
||||
require.NoError(t, err)
|
||||
stmt, err := engine.CompileToStatement(context.Background(), `sets.equivalent(tags, ["a", "b"])`, RenderOptions{Dialect: DialectSQLite})
|
||||
require.NoError(t, err)
|
||||
require.Contains(t, stmt.SQL, "JSON_ARRAY_LENGTH")
|
||||
require.Contains(t, stmt.Args, int64(2))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// exists_one() comprehension on tags
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCompileExistsOnePerDialect(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine, err := NewEngine(NewSchema())
|
||||
require.NoError(t, err)
|
||||
|
||||
cases := []struct {
|
||||
dialect DialectName
|
||||
fragments []string
|
||||
}{
|
||||
{DialectSQLite, []string{"COUNT(", "json_each(", ") = 1"}},
|
||||
{DialectPostgres, []string{"COUNT(", "jsonb_array_elements_text(", ") = 1"}},
|
||||
{DialectMySQL, []string{"COUNT(", "JSON_TABLE(", ") = 1"}},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
stmt, err := engine.CompileToStatement(context.Background(), `tags.exists_one(t, t == "urgent")`, RenderOptions{Dialect: tc.dialect})
|
||||
require.NoError(t, err, tc.dialect)
|
||||
for _, frag := range tc.fragments {
|
||||
require.Contains(t, stmt.SQL, frag, "dialect %s", tc.dialect)
|
||||
}
|
||||
require.Equal(t, []any{"urgent"}, stmt.Args, "dialect %s", tc.dialect)
|
||||
}
|
||||
}
|
||||
|
|
@ -134,6 +134,15 @@ type FunctionValue struct {
|
|||
|
||||
func (*FunctionValue) isValueExpr() {}
|
||||
|
||||
// FieldAccessorValue captures a CEL timestamp accessor on a field, such as
|
||||
// created_ts.getMonth(). It renders to a dialect-specific date-part extraction.
|
||||
type FieldAccessorValue struct {
|
||||
Field string
|
||||
Accessor string // e.g. "getFullYear", "getMonth"
|
||||
}
|
||||
|
||||
func (*FieldAccessorValue) isValueExpr() {}
|
||||
|
||||
// ListComprehensionCondition represents CEL macros like exists(), all(), filter().
|
||||
type ListComprehensionCondition struct {
|
||||
Kind ComprehensionKind
|
||||
|
|
@ -148,8 +157,9 @@ func (*ListComprehensionCondition) isCondition() {}
|
|||
type ComprehensionKind string
|
||||
|
||||
const (
|
||||
ComprehensionExists ComprehensionKind = "exists"
|
||||
ComprehensionAll ComprehensionKind = "all"
|
||||
ComprehensionExists ComprehensionKind = "exists"
|
||||
ComprehensionAll ComprehensionKind = "all"
|
||||
ComprehensionExistsOne ComprehensionKind = "exists_one"
|
||||
)
|
||||
|
||||
// PredicateExpr represents predicates used in comprehensions.
|
||||
|
|
|
|||
|
|
@ -7,10 +7,18 @@ import (
|
|||
exprv1 "google.golang.org/genproto/googleapis/api/expr/v1alpha1"
|
||||
)
|
||||
|
||||
func buildCondition(expr *exprv1.Expr, schema Schema) (Condition, error) {
|
||||
// parseContext carries the schema plus the frozen evaluation time used to fold
|
||||
// the `now` variable into a constant. Freezing once per compile guarantees a
|
||||
// single filter observes a single instant.
|
||||
type parseContext struct {
|
||||
schema Schema
|
||||
now time.Time
|
||||
}
|
||||
|
||||
func buildCondition(expr *exprv1.Expr, pc parseContext) (Condition, error) {
|
||||
switch v := expr.ExprKind.(type) {
|
||||
case *exprv1.Expr_CallExpr:
|
||||
return buildCallCondition(v.CallExpr, schema)
|
||||
return buildCallCondition(v.CallExpr, pc)
|
||||
case *exprv1.Expr_ConstExpr:
|
||||
val, err := getConstValue(expr)
|
||||
if err != nil {
|
||||
|
|
@ -22,7 +30,7 @@ func buildCondition(expr *exprv1.Expr, schema Schema) (Condition, error) {
|
|||
return nil, errors.New("filter must evaluate to a boolean value")
|
||||
case *exprv1.Expr_IdentExpr:
|
||||
name := v.IdentExpr.GetName()
|
||||
field, ok := schema.Field(name)
|
||||
field, ok := pc.schema.Field(name)
|
||||
if !ok {
|
||||
return nil, errors.Errorf("unknown identifier %q", name)
|
||||
}
|
||||
|
|
@ -31,23 +39,23 @@ func buildCondition(expr *exprv1.Expr, schema Schema) (Condition, error) {
|
|||
}
|
||||
return &FieldPredicateCondition{Field: name}, nil
|
||||
case *exprv1.Expr_ComprehensionExpr:
|
||||
return buildComprehensionCondition(v.ComprehensionExpr, schema)
|
||||
return buildComprehensionCondition(v.ComprehensionExpr, pc.schema)
|
||||
default:
|
||||
return nil, errors.New("unsupported top-level expression")
|
||||
}
|
||||
}
|
||||
|
||||
func buildCallCondition(call *exprv1.Expr_Call, schema Schema) (Condition, error) {
|
||||
func buildCallCondition(call *exprv1.Expr_Call, pc parseContext) (Condition, error) {
|
||||
switch call.Function {
|
||||
case "_&&_":
|
||||
if len(call.Args) != 2 {
|
||||
return nil, errors.New("logical AND expects two arguments")
|
||||
}
|
||||
left, err := buildCondition(call.Args[0], schema)
|
||||
left, err := buildCondition(call.Args[0], pc)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
right, err := buildCondition(call.Args[1], schema)
|
||||
right, err := buildCondition(call.Args[1], pc)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -60,11 +68,11 @@ func buildCallCondition(call *exprv1.Expr_Call, schema Schema) (Condition, error
|
|||
if len(call.Args) != 2 {
|
||||
return nil, errors.New("logical OR expects two arguments")
|
||||
}
|
||||
left, err := buildCondition(call.Args[0], schema)
|
||||
left, err := buildCondition(call.Args[0], pc)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
right, err := buildCondition(call.Args[1], schema)
|
||||
right, err := buildCondition(call.Args[1], pc)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -77,23 +85,25 @@ func buildCallCondition(call *exprv1.Expr_Call, schema Schema) (Condition, error
|
|||
if len(call.Args) != 1 {
|
||||
return nil, errors.New("logical NOT expects one argument")
|
||||
}
|
||||
child, err := buildCondition(call.Args[0], schema)
|
||||
child, err := buildCondition(call.Args[0], pc)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &NotCondition{Expr: child}, nil
|
||||
case "_==_", "_!=_", "_<_", "_>_", "_<=_", "_>=_":
|
||||
return buildComparisonCondition(call, schema)
|
||||
return buildComparisonCondition(call, pc)
|
||||
case "@in":
|
||||
return buildInCondition(call, schema)
|
||||
return buildInCondition(call, pc)
|
||||
case "contains":
|
||||
return buildTextMatchCondition(call, schema, TextMatchContains)
|
||||
return buildTextMatchCondition(call, pc.schema, TextMatchContains)
|
||||
case "startsWith":
|
||||
return buildTextMatchCondition(call, schema, TextMatchPrefix)
|
||||
return buildTextMatchCondition(call, pc.schema, TextMatchPrefix)
|
||||
case "endsWith":
|
||||
return buildTextMatchCondition(call, schema, TextMatchSuffix)
|
||||
return buildTextMatchCondition(call, pc.schema, TextMatchSuffix)
|
||||
case "matches":
|
||||
return buildMatchesCondition(call, schema)
|
||||
return buildMatchesCondition(call, pc.schema)
|
||||
case "sets.contains", "sets.intersects", "sets.equivalent":
|
||||
return buildSetCondition(call, pc)
|
||||
default:
|
||||
val, ok, err := evaluateBool(call)
|
||||
if err != nil {
|
||||
|
|
@ -106,7 +116,7 @@ func buildCallCondition(call *exprv1.Expr_Call, schema Schema) (Condition, error
|
|||
}
|
||||
}
|
||||
|
||||
func buildComparisonCondition(call *exprv1.Expr_Call, schema Schema) (Condition, error) {
|
||||
func buildComparisonCondition(call *exprv1.Expr_Call, pc parseContext) (Condition, error) {
|
||||
if len(call.Args) != 2 {
|
||||
return nil, errors.New("comparison expects two arguments")
|
||||
}
|
||||
|
|
@ -115,23 +125,23 @@ func buildComparisonCondition(call *exprv1.Expr_Call, schema Schema) (Condition,
|
|||
return nil, err
|
||||
}
|
||||
|
||||
left, err := buildValueExpr(call.Args[0], schema)
|
||||
left, err := buildValueExpr(call.Args[0], pc)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
right, err := buildValueExpr(call.Args[1], schema)
|
||||
right, err := buildValueExpr(call.Args[1], pc)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// If the left side is a field, validate allowed operators.
|
||||
if field, ok := left.(*FieldRef); ok {
|
||||
def, exists := schema.Field(field.Name)
|
||||
def, exists := pc.schema.Field(field.Name)
|
||||
if !exists {
|
||||
return nil, errors.Errorf("unknown identifier %q", field.Name)
|
||||
}
|
||||
if def.Kind == FieldKindVirtualAlias {
|
||||
def, exists = schema.ResolveAlias(field.Name)
|
||||
def, exists = pc.schema.ResolveAlias(field.Name)
|
||||
if !exists {
|
||||
return nil, errors.Errorf("invalid alias %q", field.Name)
|
||||
}
|
||||
|
|
@ -150,15 +160,15 @@ func buildComparisonCondition(call *exprv1.Expr_Call, schema Schema) (Condition,
|
|||
}, nil
|
||||
}
|
||||
|
||||
func buildInCondition(call *exprv1.Expr_Call, schema Schema) (Condition, error) {
|
||||
func buildInCondition(call *exprv1.Expr_Call, pc parseContext) (Condition, error) {
|
||||
if len(call.Args) != 2 {
|
||||
return nil, errors.New("in operator expects two arguments")
|
||||
}
|
||||
|
||||
// Handle identifier in list syntax.
|
||||
if identName, err := getIdentName(call.Args[0]); err == nil {
|
||||
if field, ok := schema.Field(identName); ok && field.Kind == FieldKindVirtualAlias {
|
||||
if _, aliasOk := schema.ResolveAlias(identName); !aliasOk {
|
||||
if field, ok := pc.schema.Field(identName); ok && field.Kind == FieldKindVirtualAlias {
|
||||
if _, aliasOk := pc.schema.ResolveAlias(identName); !aliasOk {
|
||||
return nil, errors.Errorf("invalid alias %q", identName)
|
||||
}
|
||||
} else if !ok {
|
||||
|
|
@ -168,7 +178,7 @@ func buildInCondition(call *exprv1.Expr_Call, schema Schema) (Condition, error)
|
|||
if listExpr := call.Args[1].GetListExpr(); listExpr != nil {
|
||||
values := make([]ValueExpr, 0, len(listExpr.Elements))
|
||||
for _, element := range listExpr.Elements {
|
||||
value, err := buildValueExpr(element, schema)
|
||||
value, err := buildValueExpr(element, pc)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -183,10 +193,10 @@ func buildInCondition(call *exprv1.Expr_Call, schema Schema) (Condition, error)
|
|||
|
||||
// Handle "value in identifier" syntax.
|
||||
if identName, err := getIdentName(call.Args[1]); err == nil {
|
||||
if _, ok := schema.Field(identName); !ok {
|
||||
if _, ok := pc.schema.Field(identName); !ok {
|
||||
return nil, errors.Errorf("unknown identifier %q", identName)
|
||||
}
|
||||
element, err := buildValueExpr(call.Args[0], schema)
|
||||
element, err := buildValueExpr(call.Args[0], pc)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -266,9 +276,13 @@ func buildMatchesCondition(call *exprv1.Expr_Call, schema Schema) (Condition, er
|
|||
}, nil
|
||||
}
|
||||
|
||||
func buildValueExpr(expr *exprv1.Expr, schema Schema) (ValueExpr, error) {
|
||||
func buildValueExpr(expr *exprv1.Expr, pc parseContext) (ValueExpr, error) {
|
||||
if identName, err := getIdentName(expr); err == nil {
|
||||
if _, ok := schema.Field(identName); !ok {
|
||||
// `now` is not a schema field; it folds to the frozen evaluation time.
|
||||
if identName == "now" {
|
||||
return &LiteralValue{Value: pc.now.Unix()}, nil
|
||||
}
|
||||
if _, ok := pc.schema.Field(identName); !ok {
|
||||
return nil, errors.Errorf("unknown identifier %q", identName)
|
||||
}
|
||||
return &FieldRef{Name: identName}, nil
|
||||
|
|
@ -278,7 +292,7 @@ func buildValueExpr(expr *exprv1.Expr, schema Schema) (ValueExpr, error) {
|
|||
return &LiteralValue{Value: literal}, nil
|
||||
}
|
||||
|
||||
if value, ok, err := evaluateNumeric(expr); err != nil {
|
||||
if value, ok, err := evaluateNumeric(expr, pc.now); err != nil {
|
||||
return nil, err
|
||||
} else if ok {
|
||||
return &LiteralValue{Value: value}, nil
|
||||
|
|
@ -291,12 +305,15 @@ func buildValueExpr(expr *exprv1.Expr, schema Schema) (ValueExpr, error) {
|
|||
}
|
||||
|
||||
if call := expr.GetCallExpr(); call != nil {
|
||||
if call.Target != nil && isTimestampAccessor(call.Function) {
|
||||
return buildTimestampAccessor(call, pc.schema)
|
||||
}
|
||||
switch call.Function {
|
||||
case "size":
|
||||
if len(call.Args) != 1 {
|
||||
return nil, errors.New("size() expects one argument")
|
||||
}
|
||||
arg, err := buildValueExpr(call.Args[0], schema)
|
||||
arg, err := buildValueExpr(call.Args[0], pc)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -304,10 +321,8 @@ func buildValueExpr(expr *exprv1.Expr, schema Schema) (ValueExpr, error) {
|
|||
Name: "size",
|
||||
Args: []ValueExpr{arg},
|
||||
}, nil
|
||||
case "now":
|
||||
return &LiteralValue{Value: timeNowUnix()}, nil
|
||||
case "_+_", "_-_", "_*_":
|
||||
value, ok, err := evaluateNumeric(expr)
|
||||
value, ok, err := evaluateNumeric(expr, pc.now)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -396,7 +411,11 @@ func evaluateBoolExpr(expr *exprv1.Expr) (bool, bool, error) {
|
|||
return false, false, nil
|
||||
}
|
||||
|
||||
func evaluateNumeric(expr *exprv1.Expr) (int64, bool, error) {
|
||||
// evaluateNumeric constant-folds an expression to an int64 measured in seconds:
|
||||
// timestamps and `now` fold to Unix epoch seconds, durations fold to a number of
|
||||
// seconds, and the two combine through standard arithmetic. CEL has already
|
||||
// type-checked the operand combinations, so the folded int math is well-formed.
|
||||
func evaluateNumeric(expr *exprv1.Expr, now time.Time) (int64, bool, error) {
|
||||
if literal, err := getConstValue(expr); err == nil {
|
||||
switch v := literal.(type) {
|
||||
case int64:
|
||||
|
|
@ -407,26 +426,36 @@ func evaluateNumeric(expr *exprv1.Expr) (int64, bool, error) {
|
|||
return 0, false, nil
|
||||
}
|
||||
|
||||
// The `now` variable folds to the frozen evaluation time.
|
||||
if ident := expr.GetIdentExpr(); ident != nil {
|
||||
if ident.GetName() == "now" {
|
||||
return now.Unix(), true, nil
|
||||
}
|
||||
return 0, false, nil
|
||||
}
|
||||
|
||||
call := expr.GetCallExpr()
|
||||
if call == nil {
|
||||
return 0, false, nil
|
||||
}
|
||||
|
||||
switch call.Function {
|
||||
case "now":
|
||||
return timeNowUnix(), true, nil
|
||||
case "_+_", "_-_", "_*_":
|
||||
case "timestamp":
|
||||
return evaluateTimestamp(call)
|
||||
case "duration":
|
||||
return evaluateDuration(call)
|
||||
case "_+_", "_-_", "_*_", "_/_", "_%_":
|
||||
if len(call.Args) != 2 {
|
||||
return 0, false, errors.New("arithmetic requires two arguments")
|
||||
}
|
||||
left, ok, err := evaluateNumeric(call.Args[0])
|
||||
left, ok, err := evaluateNumeric(call.Args[0], now)
|
||||
if err != nil {
|
||||
return 0, false, err
|
||||
}
|
||||
if !ok {
|
||||
return 0, false, nil
|
||||
}
|
||||
right, ok, err := evaluateNumeric(call.Args[1])
|
||||
right, ok, err := evaluateNumeric(call.Args[1], now)
|
||||
if err != nil {
|
||||
return 0, false, err
|
||||
}
|
||||
|
|
@ -440,6 +469,16 @@ func evaluateNumeric(expr *exprv1.Expr) (int64, bool, error) {
|
|||
return left - right, true, nil
|
||||
case "_*_":
|
||||
return left * right, true, nil
|
||||
case "_/_":
|
||||
if right == 0 {
|
||||
return 0, false, errors.New("division by zero")
|
||||
}
|
||||
return left / right, true, nil
|
||||
case "_%_":
|
||||
if right == 0 {
|
||||
return 0, false, errors.New("modulo by zero")
|
||||
}
|
||||
return left % right, true, nil
|
||||
default:
|
||||
return 0, false, errors.Errorf("unsupported arithmetic operator %q", call.Function)
|
||||
}
|
||||
|
|
@ -448,8 +487,185 @@ func evaluateNumeric(expr *exprv1.Expr) (int64, bool, error) {
|
|||
}
|
||||
}
|
||||
|
||||
func timeNowUnix() int64 {
|
||||
return time.Now().Unix()
|
||||
// evaluateTimestamp folds timestamp("RFC3339") and timestamp(<epoch int>) into
|
||||
// Unix epoch seconds.
|
||||
func evaluateTimestamp(call *exprv1.Expr_Call) (int64, bool, error) {
|
||||
if len(call.Args) != 1 {
|
||||
return 0, false, errors.New("timestamp() expects one argument")
|
||||
}
|
||||
value, err := getConstValue(call.Args[0])
|
||||
if err != nil {
|
||||
return 0, false, errors.Wrap(err, "timestamp() only supports literal arguments")
|
||||
}
|
||||
switch v := value.(type) {
|
||||
case string:
|
||||
ts, err := time.Parse(time.RFC3339, v)
|
||||
if err != nil {
|
||||
return 0, false, errors.Wrap(err, "invalid timestamp literal")
|
||||
}
|
||||
return ts.Unix(), true, nil
|
||||
case int64:
|
||||
return v, true, nil
|
||||
default:
|
||||
return 0, false, errors.New("timestamp() argument must be an RFC3339 string or epoch int")
|
||||
}
|
||||
}
|
||||
|
||||
// evaluateDuration folds duration("<go-duration>") into a number of seconds.
|
||||
func evaluateDuration(call *exprv1.Expr_Call) (int64, bool, error) {
|
||||
if len(call.Args) != 1 {
|
||||
return 0, false, errors.New("duration() expects one argument")
|
||||
}
|
||||
value, err := getConstValue(call.Args[0])
|
||||
if err != nil {
|
||||
return 0, false, errors.Wrap(err, "duration() only supports literal arguments")
|
||||
}
|
||||
str, ok := value.(string)
|
||||
if !ok {
|
||||
return 0, false, errors.New("duration() argument must be a string")
|
||||
}
|
||||
d, err := time.ParseDuration(str)
|
||||
if err != nil {
|
||||
return 0, false, errors.Wrap(err, "invalid duration literal")
|
||||
}
|
||||
return int64(d.Seconds()), true, nil
|
||||
}
|
||||
|
||||
// timestampAccessors is the set of supported CEL timestamp accessor methods.
|
||||
var timestampAccessors = map[string]bool{
|
||||
"getFullYear": true,
|
||||
"getMonth": true,
|
||||
"getDate": true,
|
||||
"getDayOfMonth": true,
|
||||
"getDayOfWeek": true,
|
||||
"getDayOfYear": true,
|
||||
"getHours": true,
|
||||
"getMinutes": true,
|
||||
"getSeconds": true,
|
||||
}
|
||||
|
||||
func isTimestampAccessor(name string) bool {
|
||||
return timestampAccessors[name]
|
||||
}
|
||||
|
||||
// buildTimestampAccessor converts created_ts.getMonth() into a FieldAccessorValue.
|
||||
// Timezone arguments are rejected; extraction is UTC (see renderer).
|
||||
func buildTimestampAccessor(call *exprv1.Expr_Call, schema Schema) (ValueExpr, error) {
|
||||
targetName, err := getIdentName(call.Target)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "timestamp accessor requires a field target")
|
||||
}
|
||||
field, ok := schema.Field(targetName)
|
||||
if !ok {
|
||||
return nil, errors.Errorf("unknown identifier %q", targetName)
|
||||
}
|
||||
if field.Type != FieldTypeTimestamp {
|
||||
return nil, errors.Errorf("%s() is only valid on timestamp fields, got %q", call.Function, targetName)
|
||||
}
|
||||
if len(call.Args) != 0 {
|
||||
return nil, errors.Errorf("%s() with a timezone argument is not supported", call.Function)
|
||||
}
|
||||
return &FieldAccessorValue{Field: targetName, Accessor: call.Function}, nil
|
||||
}
|
||||
|
||||
// buildSetCondition desugars ext.Sets() operations over a JSON list field into
|
||||
// existing IR: membership reduces to ElementInCondition, and equivalence adds a
|
||||
// length check. This relies on the list field being a set (no duplicates), which
|
||||
// holds for memo tags.
|
||||
func buildSetCondition(call *exprv1.Expr_Call, pc parseContext) (Condition, error) {
|
||||
if len(call.Args) != 2 {
|
||||
return nil, errors.Errorf("%s expects two arguments", call.Function)
|
||||
}
|
||||
|
||||
fieldName, err := getIdentName(call.Args[0])
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "set operations require a list field as the first argument")
|
||||
}
|
||||
field, ok := pc.schema.Field(fieldName)
|
||||
if !ok {
|
||||
return nil, errors.Errorf("unknown identifier %q", fieldName)
|
||||
}
|
||||
if field.Kind != FieldKindJSONList {
|
||||
return nil, errors.Errorf("set operations require a list field, got %q", fieldName)
|
||||
}
|
||||
|
||||
listExpr := call.Args[1].GetListExpr()
|
||||
if listExpr == nil {
|
||||
return nil, errors.New("set operations require a list literal as the second argument")
|
||||
}
|
||||
values := make([]string, 0, len(listExpr.Elements))
|
||||
for _, el := range listExpr.Elements {
|
||||
v, err := getConstValue(el)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "set operations only support literal string elements")
|
||||
}
|
||||
s, ok := v.(string)
|
||||
if !ok {
|
||||
return nil, errors.New("set operations require string elements")
|
||||
}
|
||||
values = append(values, s)
|
||||
}
|
||||
|
||||
membership := func(s string) Condition {
|
||||
return &ElementInCondition{Element: &LiteralValue{Value: s}, Field: fieldName}
|
||||
}
|
||||
sizeEquals := func(n int) Condition {
|
||||
return &ComparisonCondition{
|
||||
Left: &FunctionValue{Name: "size", Args: []ValueExpr{&FieldRef{Name: fieldName}}},
|
||||
Operator: CompareEq,
|
||||
Right: &LiteralValue{Value: int64(n)},
|
||||
}
|
||||
}
|
||||
|
||||
switch call.Function {
|
||||
case "sets.contains":
|
||||
if len(values) == 0 {
|
||||
return &ConstantCondition{Value: true}, nil
|
||||
}
|
||||
return combineConditions(LogicalAnd, mapConditions(values, membership)), nil
|
||||
case "sets.intersects":
|
||||
if len(values) == 0 {
|
||||
return &ConstantCondition{Value: false}, nil
|
||||
}
|
||||
return combineConditions(LogicalOr, mapConditions(values, membership)), nil
|
||||
case "sets.equivalent":
|
||||
distinct := distinctStrings(values)
|
||||
if len(distinct) == 0 {
|
||||
return sizeEquals(0), nil
|
||||
}
|
||||
contains := combineConditions(LogicalAnd, mapConditions(distinct, membership))
|
||||
return &LogicalCondition{Operator: LogicalAnd, Left: contains, Right: sizeEquals(len(distinct))}, nil
|
||||
default:
|
||||
return nil, errors.Errorf("unsupported set operation %q", call.Function)
|
||||
}
|
||||
}
|
||||
|
||||
func mapConditions(values []string, f func(string) Condition) []Condition {
|
||||
conds := make([]Condition, 0, len(values))
|
||||
for _, v := range values {
|
||||
conds = append(conds, f(v))
|
||||
}
|
||||
return conds
|
||||
}
|
||||
|
||||
func combineConditions(op LogicalOperator, conds []Condition) Condition {
|
||||
result := conds[0]
|
||||
for _, c := range conds[1:] {
|
||||
result = &LogicalCondition{Operator: op, Left: result, Right: c}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func distinctStrings(values []string) []string {
|
||||
seen := make(map[string]bool, len(values))
|
||||
out := make([]string, 0, len(values))
|
||||
for _, v := range values {
|
||||
if !seen[v] {
|
||||
seen[v] = true
|
||||
out = append(out, v)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// buildComprehensionCondition handles CEL comprehension expressions (exists, all, etc.).
|
||||
|
|
@ -513,7 +729,15 @@ func detectComprehensionKind(comp *exprv1.Expr_Comprehension) (ComprehensionKind
|
|||
}
|
||||
}
|
||||
|
||||
return "", errors.New("unsupported comprehension type; only exists() is supported")
|
||||
// exists_one() starts at int(0) and increments via a conditional (predicate ?
|
||||
// accu + 1 : accu) in the loop step.
|
||||
if _, isInt := accuInit.GetConstantKind().(*exprv1.Constant_Int64Value); isInt {
|
||||
if step := comp.LoopStep.GetCallExpr(); step != nil && step.Function == "_?_:_" {
|
||||
return ComprehensionExistsOne, nil
|
||||
}
|
||||
}
|
||||
|
||||
return "", errors.New("unsupported comprehension type (supported: exists, all, exists_one)")
|
||||
}
|
||||
|
||||
// extractPredicate extracts the predicate expression from the comprehension loop step.
|
||||
|
|
@ -525,12 +749,20 @@ func extractPredicate(comp *exprv1.Expr_Comprehension, _ Schema) (PredicateExpr,
|
|||
return nil, errors.New("comprehension loop step must be a call expression")
|
||||
}
|
||||
|
||||
if len(step.Args) != 2 {
|
||||
return nil, errors.New("comprehension loop step must have two arguments")
|
||||
// exists/all: accu || predicate / accu && predicate -> predicate is arg[1].
|
||||
// exists_one: predicate ? accu + 1 : accu -> predicate is arg[0].
|
||||
var predicateExpr *exprv1.Expr
|
||||
if step.Function == "_?_:_" {
|
||||
if len(step.Args) != 3 {
|
||||
return nil, errors.New("exists_one loop step must have three arguments")
|
||||
}
|
||||
predicateExpr = step.Args[0]
|
||||
} else {
|
||||
if len(step.Args) != 2 {
|
||||
return nil, errors.New("comprehension loop step must have two arguments")
|
||||
}
|
||||
predicateExpr = step.Args[1]
|
||||
}
|
||||
|
||||
// The predicate is the second argument
|
||||
predicateExpr := step.Args[1]
|
||||
predicateCall := predicateExpr.GetCallExpr()
|
||||
if predicateCall == nil {
|
||||
return nil, errors.New("comprehension predicate must be a function call")
|
||||
|
|
|
|||
|
|
@ -167,11 +167,85 @@ func (r *renderer) renderComparison(cond *ComparisonCondition) (renderResult, er
|
|||
}
|
||||
case *FunctionValue:
|
||||
return r.renderFunctionComparison(left, cond.Operator, cond.Right)
|
||||
case *FieldAccessorValue:
|
||||
return r.renderAccessorComparison(left, cond.Operator, cond.Right)
|
||||
default:
|
||||
return renderResult{}, errors.New("comparison must start with a field reference or supported function")
|
||||
}
|
||||
}
|
||||
|
||||
// accessorSpec maps a CEL timestamp accessor to per-dialect SQL date-part tokens
|
||||
// and the offset to subtract so the result matches CEL's base (e.g. CEL months
|
||||
// are 0-based but every dialect reports 1-based, so off=1). off is indexed
|
||||
// [sqlite, postgres, mysql].
|
||||
type accessorSpec struct {
|
||||
sqlite string // strftime format specifier
|
||||
pg string // EXTRACT field
|
||||
mysql string // function name
|
||||
off [3]int
|
||||
}
|
||||
|
||||
var accessorSpecs = map[string]accessorSpec{
|
||||
"getFullYear": {"%Y", "YEAR", "YEAR", [3]int{0, 0, 0}},
|
||||
"getMonth": {"%m", "MONTH", "MONTH", [3]int{1, 1, 1}},
|
||||
"getDate": {"%d", "DAY", "DAYOFMONTH", [3]int{0, 0, 0}},
|
||||
"getDayOfMonth": {"%d", "DAY", "DAYOFMONTH", [3]int{1, 1, 1}},
|
||||
"getDayOfWeek": {"%w", "DOW", "DAYOFWEEK", [3]int{0, 0, 1}},
|
||||
"getDayOfYear": {"%j", "DOY", "DAYOFYEAR", [3]int{1, 1, 1}},
|
||||
"getHours": {"%H", "HOUR", "HOUR", [3]int{0, 0, 0}},
|
||||
"getMinutes": {"%M", "MINUTE", "MINUTE", [3]int{0, 0, 0}},
|
||||
"getSeconds": {"%S", "SECOND", "SECOND", [3]int{0, 0, 0}},
|
||||
}
|
||||
|
||||
func (r *renderer) renderAccessorComparison(acc *FieldAccessorValue, op ComparisonOperator, right ValueExpr) (renderResult, error) {
|
||||
field, ok := r.schema.Field(acc.Field)
|
||||
if !ok {
|
||||
return renderResult{}, errors.Errorf("unknown field %q", acc.Field)
|
||||
}
|
||||
value, err := expectNumericLiteral(right)
|
||||
if err != nil {
|
||||
return renderResult{}, err
|
||||
}
|
||||
expr, err := r.timestampAccessorExpr(field, acc.Accessor)
|
||||
if err != nil {
|
||||
return renderResult{}, err
|
||||
}
|
||||
placeholder := r.addArg(value)
|
||||
return renderResult{
|
||||
sql: fmt.Sprintf("%s %s %s", expr, sqlOperator(op), placeholder),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// timestampAccessorExpr builds a dialect-specific integer expression for a CEL
|
||||
// timestamp accessor. Extraction is UTC on SQLite/Postgres (epoch columns); on
|
||||
// MySQL the TIMESTAMP column is read in the server session time zone.
|
||||
func (r *renderer) timestampAccessorExpr(field Field, accessor string) (string, error) {
|
||||
spec, ok := accessorSpecs[accessor]
|
||||
if !ok {
|
||||
return "", errors.Errorf("unsupported timestamp accessor %q", accessor)
|
||||
}
|
||||
col := qualifyColumn(r.dialect, field.Column)
|
||||
var base string
|
||||
var off int
|
||||
switch r.dialect {
|
||||
case DialectSQLite:
|
||||
base = fmt.Sprintf("CAST(strftime('%s', %s, 'unixepoch') AS INTEGER)", spec.sqlite, col)
|
||||
off = spec.off[0]
|
||||
case DialectPostgres:
|
||||
base = fmt.Sprintf("CAST(EXTRACT(%s FROM to_timestamp(%s) AT TIME ZONE 'UTC') AS INTEGER)", spec.pg, col)
|
||||
off = spec.off[1]
|
||||
case DialectMySQL:
|
||||
base = fmt.Sprintf("%s(%s)", spec.mysql, col)
|
||||
off = spec.off[2]
|
||||
default:
|
||||
return "", errors.Errorf("unsupported dialect %q", r.dialect)
|
||||
}
|
||||
if off != 0 {
|
||||
base = fmt.Sprintf("(%s - %d)", base, off)
|
||||
}
|
||||
return base, nil
|
||||
}
|
||||
|
||||
func (r *renderer) renderFunctionComparison(fn *FunctionValue, op ComparisonOperator, right ValueExpr) (renderResult, error) {
|
||||
if fn.Name != "size" {
|
||||
return renderResult{}, errors.Errorf("unsupported function %s in comparison", fn.Name)
|
||||
|
|
@ -188,8 +262,11 @@ func (r *renderer) renderFunctionComparison(fn *FunctionValue, op ComparisonOper
|
|||
if !ok {
|
||||
return renderResult{}, errors.Errorf("unknown field %q", fieldArg.Name)
|
||||
}
|
||||
if field.Kind != FieldKindJSONList {
|
||||
return renderResult{}, errors.Errorf("size() only supports tag lists, got %q", field.Name)
|
||||
if field.Kind == FieldKindVirtualAlias {
|
||||
field, ok = r.schema.ResolveAlias(fieldArg.Name)
|
||||
if !ok {
|
||||
return renderResult{}, errors.Errorf("invalid alias %q", fieldArg.Name)
|
||||
}
|
||||
}
|
||||
|
||||
value, err := expectNumericLiteral(right)
|
||||
|
|
@ -197,13 +274,33 @@ func (r *renderer) renderFunctionComparison(fn *FunctionValue, op ComparisonOper
|
|||
return renderResult{}, err
|
||||
}
|
||||
|
||||
expr := jsonArrayLengthExpr(r.dialect, field)
|
||||
var expr string
|
||||
switch {
|
||||
case field.Kind == FieldKindJSONList:
|
||||
expr = jsonArrayLengthExpr(r.dialect, field)
|
||||
case field.Kind == FieldKindScalar && field.Type == FieldTypeString:
|
||||
expr = stringLengthExpr(r.dialect, field.columnExpr(r.dialect))
|
||||
default:
|
||||
return renderResult{}, errors.Errorf("size() does not support field %q", field.Name)
|
||||
}
|
||||
|
||||
placeholder := r.addArg(value)
|
||||
return renderResult{
|
||||
sql: fmt.Sprintf("%s %s %s", expr, sqlOperator(op), placeholder),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// stringLengthExpr returns the character-count expression for a string column.
|
||||
// MySQL's LENGTH counts bytes, so CHAR_LENGTH is used to count characters and
|
||||
// match CEL's size() code-point semantics; SQLite/Postgres LENGTH already counts
|
||||
// characters.
|
||||
func stringLengthExpr(d DialectName, colExpr string) string {
|
||||
if d == DialectMySQL {
|
||||
return fmt.Sprintf("CHAR_LENGTH(%s)", colExpr)
|
||||
}
|
||||
return fmt.Sprintf("LENGTH(%s)", colExpr)
|
||||
}
|
||||
|
||||
func (r *renderer) renderScalarComparison(field Field, op ComparisonOperator, right ValueExpr) (renderResult, error) {
|
||||
lit, err := expectLiteral(right)
|
||||
if err != nil {
|
||||
|
|
@ -525,6 +622,10 @@ func (r *renderer) renderListComprehension(cond *ListComprehensionCondition) (re
|
|||
return r.renderTagAll(field, cond.Predicate)
|
||||
}
|
||||
|
||||
if cond.Kind == ComprehensionExistsOne {
|
||||
return r.renderTagExistsOne(field, cond.Predicate)
|
||||
}
|
||||
|
||||
// Render based on predicate type
|
||||
switch pred := cond.Predicate.(type) {
|
||||
case *EqualsPredicate:
|
||||
|
|
@ -567,6 +668,27 @@ func (r *renderer) renderTagAll(field Field, pred PredicateExpr) (renderResult,
|
|||
}
|
||||
}
|
||||
|
||||
// renderTagExistsOne renders tags.exists_one(t, <pred>): exactly one element
|
||||
// satisfies the predicate, via a COUNT(...) = 1 subquery. A null or empty array
|
||||
// yields COUNT 0, which is correctly not equal to 1.
|
||||
func (r *renderer) renderTagExistsOne(field Field, pred PredicateExpr) (renderResult, error) {
|
||||
arrayExpr := jsonArrayExpr(r.dialect, field)
|
||||
elemCond, err := r.elementPredicateSQL(pred)
|
||||
if err != nil {
|
||||
return renderResult{}, err
|
||||
}
|
||||
switch r.dialect {
|
||||
case DialectSQLite:
|
||||
return renderResult{sql: fmt.Sprintf("(SELECT COUNT(*) FROM json_each(%s) WHERE %s) = 1", arrayExpr, elemCond)}, nil
|
||||
case DialectMySQL:
|
||||
return renderResult{sql: fmt.Sprintf("(SELECT COUNT(*) FROM JSON_TABLE(%s, '$[*]' COLUMNS (value VARCHAR(512) PATH '$')) AS elem WHERE %s) = 1", arrayExpr, elemCond)}, nil
|
||||
case DialectPostgres:
|
||||
return renderResult{sql: fmt.Sprintf("(SELECT COUNT(*) FROM jsonb_array_elements_text(%s) AS elem(value) WHERE %s) = 1", arrayExpr, elemCond)}, nil
|
||||
default:
|
||||
return renderResult{}, errors.Errorf("unsupported dialect %s", r.dialect)
|
||||
}
|
||||
}
|
||||
|
||||
// elementPredicateSQL builds the per-element SQL condition for an all() predicate.
|
||||
// The iterated element is exposed as the unqualified column `value` on all dialects
|
||||
// (json_each.value / JSON_TABLE column / elem(value)).
|
||||
|
|
|
|||
|
|
@ -2,11 +2,9 @@ package filter
|
|||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/google/cel-go/cel"
|
||||
"github.com/google/cel-go/common/types"
|
||||
"github.com/google/cel-go/common/types/ref"
|
||||
"github.com/google/cel-go/ext"
|
||||
)
|
||||
|
||||
// DialectName enumerates supported SQL dialects.
|
||||
|
|
@ -87,16 +85,6 @@ func (s Schema) ResolveAlias(name string) (Field, bool) {
|
|||
return field, true
|
||||
}
|
||||
|
||||
var nowFunction = cel.Function("now",
|
||||
cel.Overload("now",
|
||||
[]*cel.Type{},
|
||||
cel.IntType,
|
||||
cel.FunctionBinding(func(_ ...ref.Val) ref.Val {
|
||||
return types.Int(time.Now().Unix())
|
||||
}),
|
||||
),
|
||||
)
|
||||
|
||||
// NewSchema constructs the memo filter schema and CEL environment.
|
||||
func NewSchema() Schema {
|
||||
fields := map[string]Field{
|
||||
|
|
@ -245,8 +233,8 @@ func NewSchema() Schema {
|
|||
cel.Variable("content", cel.StringType),
|
||||
cel.Variable("creator", cel.StringType),
|
||||
cel.Variable("creator_id", cel.IntType),
|
||||
cel.Variable("created_ts", cel.IntType),
|
||||
cel.Variable("updated_ts", cel.IntType),
|
||||
cel.Variable("created_ts", cel.TimestampType),
|
||||
cel.Variable("updated_ts", cel.TimestampType),
|
||||
cel.Variable("pinned", cel.BoolType),
|
||||
cel.Variable("tag", cel.StringType),
|
||||
cel.Variable("tags", cel.ListType(cel.StringType)),
|
||||
|
|
@ -255,7 +243,8 @@ func NewSchema() Schema {
|
|||
cel.Variable("has_link", cel.BoolType),
|
||||
cel.Variable("has_code", cel.BoolType),
|
||||
cel.Variable("has_incomplete_tasks", cel.BoolType),
|
||||
nowFunction,
|
||||
cel.Variable("now", cel.TimestampType),
|
||||
ext.Sets(),
|
||||
cel.ASTValidators(cel.ValidateRegexLiterals()),
|
||||
}
|
||||
|
||||
|
|
@ -314,9 +303,9 @@ func NewAttachmentSchema() Schema {
|
|||
envOptions := []cel.EnvOption{
|
||||
cel.Variable("filename", cel.StringType),
|
||||
cel.Variable("mime_type", cel.StringType),
|
||||
cel.Variable("create_time", cel.IntType),
|
||||
cel.Variable("create_time", cel.TimestampType),
|
||||
cel.Variable("memo_id", cel.AnyType),
|
||||
nowFunction,
|
||||
cel.Variable("now", cel.TimestampType),
|
||||
cel.ASTValidators(cel.ValidateRegexLiterals()),
|
||||
}
|
||||
|
||||
|
|
|
|||
100
internal/filter/time_test.go
Normal file
100
internal/filter/time_test.go
Normal file
|
|
@ -0,0 +1,100 @@
|
|||
package filter
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// fixedClock returns a deterministic clock for asserting folded `now` values.
|
||||
func fixedClock(epoch int64) func() time.Time {
|
||||
return func() time.Time { return time.Unix(epoch, 0) }
|
||||
}
|
||||
|
||||
func memoEngineAt(t *testing.T, epoch int64) *Engine {
|
||||
t.Helper()
|
||||
engine, err := NewEngine(NewSchema())
|
||||
require.NoError(t, err)
|
||||
engine.nowFunc = fixedClock(epoch)
|
||||
return engine
|
||||
}
|
||||
|
||||
func TestCompileNowVariableFoldsToInjectedClock(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine := memoEngineAt(t, 1750000000)
|
||||
stmt, err := engine.CompileToStatement(context.Background(), `created_ts >= now`, RenderOptions{Dialect: DialectSQLite})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []any{int64(1750000000)}, stmt.Args)
|
||||
}
|
||||
|
||||
func TestCompileNowFunctionIsRemoved(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine, err := NewEngine(NewSchema())
|
||||
require.NoError(t, err)
|
||||
|
||||
// now() was the legacy custom function; it is replaced by the `now` variable.
|
||||
_, err = engine.Compile(context.Background(), `created_ts >= now()`)
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
func TestCompileNowMinusDurationFolds(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine := memoEngineAt(t, 1750000000)
|
||||
stmt, err := engine.CompileToStatement(context.Background(), `created_ts >= now - duration("1h")`, RenderOptions{Dialect: DialectSQLite})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []any{int64(1750000000 - 3600)}, stmt.Args)
|
||||
}
|
||||
|
||||
func TestCompileNowPlusDurationFolds(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine := memoEngineAt(t, 1750000000)
|
||||
stmt, err := engine.CompileToStatement(context.Background(), `updated_ts < now + duration("24h")`, RenderOptions{Dialect: DialectSQLite})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []any{int64(1750000000 + 86400)}, stmt.Args)
|
||||
}
|
||||
|
||||
func TestCompileAbsoluteTimestampStringFolds(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine := memoEngineAt(t, 1750000000)
|
||||
stmt, err := engine.CompileToStatement(context.Background(), `created_ts >= timestamp("2025-01-01T00:00:00Z")`, RenderOptions{Dialect: DialectSQLite})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []any{int64(1735689600)}, stmt.Args)
|
||||
}
|
||||
|
||||
func TestCompileTimestampFromEpochIntFolds(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// This is the shape the frontend date-range filter emits.
|
||||
engine := memoEngineAt(t, 1750000000)
|
||||
stmt, err := engine.CompileToStatement(context.Background(), `created_ts >= timestamp(1730000000)`, RenderOptions{Dialect: DialectSQLite})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []any{int64(1730000000)}, stmt.Args)
|
||||
}
|
||||
|
||||
func TestCompileInvalidDurationLiteralErrors(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine := memoEngineAt(t, 1750000000)
|
||||
_, err := engine.Compile(context.Background(), `created_ts >= now - duration("garbage")`)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "duration")
|
||||
}
|
||||
|
||||
func TestCompileAttachmentCreateTimeUsesNow(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
engine, err := NewEngine(NewAttachmentSchema())
|
||||
require.NoError(t, err)
|
||||
engine.nowFunc = fixedClock(1750000000)
|
||||
|
||||
stmt, err := engine.CompileToStatement(context.Background(), `create_time >= now - duration("24h")`, RenderOptions{Dialect: DialectSQLite})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []any{int64(1750000000 - 86400)}, stmt.Args)
|
||||
}
|
||||
|
|
@ -159,7 +159,7 @@ func TestAttachmentFilterMimeTypeInList(t *testing.T) {
|
|||
// =============================================================================
|
||||
// Create Time Field Tests
|
||||
// Schema: create_time (timestamp, all comparison operators)
|
||||
// Functions: now(), arithmetic (+, -, *)
|
||||
// Time helpers: now variable, timestamp(...), duration(...)
|
||||
// =============================================================================
|
||||
|
||||
func TestAttachmentFilterCreateTimeComparison(t *testing.T) {
|
||||
|
|
@ -171,15 +171,15 @@ func TestAttachmentFilterCreateTimeComparison(t *testing.T) {
|
|||
tc.CreateAttachment(NewAttachmentBuilder(tc.CreatorID).Filename("test.png").MimeType("image/png"))
|
||||
|
||||
// Test: create_time < future (should match)
|
||||
attachments := tc.ListWithFilter(`create_time < ` + formatInt64(now+3600))
|
||||
attachments := tc.ListWithFilter(`create_time < timestamp(` + formatInt64(now+3600) + `)`)
|
||||
require.Len(t, attachments, 1)
|
||||
|
||||
// Test: create_time > past (should match)
|
||||
attachments = tc.ListWithFilter(`create_time > ` + formatInt64(now-3600))
|
||||
attachments = tc.ListWithFilter(`create_time > timestamp(` + formatInt64(now-3600) + `)`)
|
||||
require.Len(t, attachments, 1)
|
||||
|
||||
// Test: create_time > future (should not match)
|
||||
attachments = tc.ListWithFilter(`create_time > ` + formatInt64(now+3600))
|
||||
attachments = tc.ListWithFilter(`create_time > timestamp(` + formatInt64(now+3600) + `)`)
|
||||
require.Len(t, attachments, 0)
|
||||
}
|
||||
|
||||
|
|
@ -190,12 +190,12 @@ func TestAttachmentFilterCreateTimeWithNow(t *testing.T) {
|
|||
|
||||
tc.CreateAttachment(NewAttachmentBuilder(tc.CreatorID).Filename("test.png").MimeType("image/png"))
|
||||
|
||||
// Test: create_time < now() + 5 (buffer for container clock drift)
|
||||
attachments := tc.ListWithFilter(`create_time < now() + 5`)
|
||||
// Test: create_time < now + 5s (buffer for container clock drift)
|
||||
attachments := tc.ListWithFilter(`create_time < now + duration("5s")`)
|
||||
require.Len(t, attachments, 1)
|
||||
|
||||
// Test: create_time > now() + 5 (should not match)
|
||||
attachments = tc.ListWithFilter(`create_time > now() + 5`)
|
||||
// Test: create_time > now + 5s (should not match)
|
||||
attachments = tc.ListWithFilter(`create_time > now + duration("5s")`)
|
||||
require.Len(t, attachments, 0)
|
||||
}
|
||||
|
||||
|
|
@ -206,16 +206,16 @@ func TestAttachmentFilterCreateTimeArithmetic(t *testing.T) {
|
|||
|
||||
tc.CreateAttachment(NewAttachmentBuilder(tc.CreatorID).Filename("test.png").MimeType("image/png"))
|
||||
|
||||
// Test: create_time >= now() - 3600 (attachments created in last hour)
|
||||
attachments := tc.ListWithFilter(`create_time >= now() - 3600`)
|
||||
// Test: create_time >= now - 1h (attachments created in last hour)
|
||||
attachments := tc.ListWithFilter(`create_time >= now - duration("1h")`)
|
||||
require.Len(t, attachments, 1)
|
||||
|
||||
// Test: create_time < now() - 86400 (attachments older than 1 day - should be empty)
|
||||
attachments = tc.ListWithFilter(`create_time < now() - 86400`)
|
||||
// Test: create_time < now - 24h (attachments older than 1 day - should be empty)
|
||||
attachments = tc.ListWithFilter(`create_time < now - duration("24h")`)
|
||||
require.Len(t, attachments, 0)
|
||||
|
||||
// Test: Multiplication - create_time >= now() - 60 * 60
|
||||
attachments = tc.ListWithFilter(`create_time >= now() - 60 * 60`)
|
||||
// Test: chained duration arithmetic - create_time >= now - 30m - 30m
|
||||
attachments = tc.ListWithFilter(`create_time >= now - duration("30m") - duration("30m")`)
|
||||
require.Len(t, attachments, 1)
|
||||
}
|
||||
|
||||
|
|
@ -227,19 +227,19 @@ func TestAttachmentFilterAllComparisonOperators(t *testing.T) {
|
|||
tc.CreateAttachment(NewAttachmentBuilder(tc.CreatorID).Filename("test.png").MimeType("image/png"))
|
||||
|
||||
// Test: < (less than)
|
||||
attachments := tc.ListWithFilter(`create_time < now() + 3600`)
|
||||
attachments := tc.ListWithFilter(`create_time < now + duration("1h")`)
|
||||
require.Len(t, attachments, 1)
|
||||
|
||||
// Test: <= (less than or equal) with buffer for clock drift
|
||||
attachments = tc.ListWithFilter(`create_time < now() + 5`)
|
||||
attachments = tc.ListWithFilter(`create_time < now + duration("5s")`)
|
||||
require.Len(t, attachments, 1)
|
||||
|
||||
// Test: > (greater than)
|
||||
attachments = tc.ListWithFilter(`create_time > now() - 3600`)
|
||||
attachments = tc.ListWithFilter(`create_time > now - duration("1h")`)
|
||||
require.Len(t, attachments, 1)
|
||||
|
||||
// Test: >= (greater than or equal)
|
||||
attachments = tc.ListWithFilter(`create_time >= now() - 60`)
|
||||
attachments = tc.ListWithFilter(`create_time >= now - duration("60s")`)
|
||||
require.Len(t, attachments, 1)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -66,6 +66,11 @@ func (b *MemoBuilder) Visibility(v store.Visibility) *MemoBuilder {
|
|||
return b
|
||||
}
|
||||
|
||||
func (b *MemoBuilder) CreatedTs(ts int64) *MemoBuilder {
|
||||
b.memo.CreatedTs = ts
|
||||
return b
|
||||
}
|
||||
|
||||
func (b *MemoBuilder) Tags(tags ...string) *MemoBuilder {
|
||||
if b.memo.Payload == nil {
|
||||
b.memo.Payload = &storepb.MemoPayload{}
|
||||
|
|
|
|||
91
store/test/memo_filter_functions_test.go
Normal file
91
store/test/memo_filter_functions_test.go
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
package test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// =============================================================================
|
||||
// size() on string fields
|
||||
// =============================================================================
|
||||
|
||||
func TestMemoFilterSizeContent(t *testing.T) {
|
||||
t.Parallel()
|
||||
tc := NewMemoFilterTestContext(t)
|
||||
defer tc.Close()
|
||||
|
||||
tc.CreateMemo(NewMemoBuilder("memo-short", tc.User.ID).Content("abc")) // length 3
|
||||
tc.CreateMemo(NewMemoBuilder("memo-long", tc.User.ID).Content("abcdefghij")) // length 10
|
||||
|
||||
require.Len(t, tc.ListWithFilter(`size(content) > 5`), 1)
|
||||
require.Len(t, tc.ListWithFilter(`size(content) < 5`), 1)
|
||||
require.Len(t, tc.ListWithFilter(`size(content) == 3`), 1)
|
||||
require.Len(t, tc.ListWithFilter(`size(content) >= 3`), 2)
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Timestamp accessors (UTC). 1700000000 == 2023-11-14 22:13:20 UTC (a Tuesday).
|
||||
// =============================================================================
|
||||
|
||||
func TestMemoFilterTimestampAccessors(t *testing.T) {
|
||||
t.Parallel()
|
||||
tc := NewMemoFilterTestContext(t)
|
||||
defer tc.Close()
|
||||
|
||||
tc.CreateMemo(NewMemoBuilder("memo-dated", tc.User.ID).Content("dated").CreatedTs(1700000000))
|
||||
|
||||
require.Len(t, tc.ListWithFilter(`created_ts.getFullYear() == 2023`), 1)
|
||||
require.Len(t, tc.ListWithFilter(`created_ts.getFullYear() == 2022`), 0)
|
||||
// getMonth is 0-based: November == 10.
|
||||
require.Len(t, tc.ListWithFilter(`created_ts.getMonth() == 10`), 1)
|
||||
require.Len(t, tc.ListWithFilter(`created_ts.getMonth() == 11`), 0)
|
||||
require.Len(t, tc.ListWithFilter(`created_ts.getDate() == 14`), 1)
|
||||
require.Len(t, tc.ListWithFilter(`created_ts.getHours() == 22`), 1)
|
||||
// getDayOfWeek is 0-based, 0=Sunday: 2023-11-14 is a Tuesday == 2.
|
||||
require.Len(t, tc.ListWithFilter(`created_ts.getDayOfWeek() == 2`), 1)
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// ext.Sets(): contains / intersects / equivalent over tags
|
||||
// =============================================================================
|
||||
|
||||
func TestMemoFilterSets(t *testing.T) {
|
||||
t.Parallel()
|
||||
tc := NewMemoFilterTestContext(t)
|
||||
defer tc.Close()
|
||||
|
||||
tc.CreateMemo(NewMemoBuilder("memo-wu", tc.User.ID).Content("work urgent").Tags("work", "urgent"))
|
||||
tc.CreateMemo(NewMemoBuilder("memo-home", tc.User.ID).Content("home").Tags("home"))
|
||||
|
||||
// contains: tags must be a superset of the list.
|
||||
require.Len(t, tc.ListWithFilter(`sets.contains(tags, ["work", "urgent"])`), 1)
|
||||
require.Len(t, tc.ListWithFilter(`sets.contains(tags, ["work", "home"])`), 0)
|
||||
|
||||
// intersects: any element in common.
|
||||
require.Len(t, tc.ListWithFilter(`sets.intersects(tags, ["urgent", "home"])`), 2)
|
||||
require.Len(t, tc.ListWithFilter(`sets.intersects(tags, ["missing"])`), 0)
|
||||
|
||||
// equivalent: exactly the same set.
|
||||
require.Len(t, tc.ListWithFilter(`sets.equivalent(tags, ["work", "urgent"])`), 1)
|
||||
require.Len(t, tc.ListWithFilter(`sets.equivalent(tags, ["work"])`), 0)
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// exists_one(): exactly one element matches the predicate
|
||||
// =============================================================================
|
||||
|
||||
func TestMemoFilterExistsOne(t *testing.T) {
|
||||
t.Parallel()
|
||||
tc := NewMemoFilterTestContext(t)
|
||||
defer tc.Close()
|
||||
|
||||
tc.CreateMemo(NewMemoBuilder("memo-one-x", tc.User.ID).Content("one x").Tags("x")) // 1 tag starts with x
|
||||
tc.CreateMemo(NewMemoBuilder("memo-two-x", tc.User.ID).Content("two x").Tags("x", "xy")) // 2 tags start with x
|
||||
tc.CreateMemo(NewMemoBuilder("memo-no-x", tc.User.ID).Content("no x").Tags("z")) // 0 tags start with x
|
||||
|
||||
// Only memo-one-x has exactly one tag starting with "x".
|
||||
require.Len(t, tc.ListWithFilter(`tags.exists_one(t, t.startsWith("x"))`), 1)
|
||||
// Only memo-no-x has exactly one tag equal to "z".
|
||||
require.Len(t, tc.ListWithFilter(`tags.exists_one(t, t == "z")`), 1)
|
||||
}
|
||||
|
|
@ -625,7 +625,7 @@ func TestMemoFilterCombinedJSONBool(t *testing.T) {
|
|||
// =============================================================================
|
||||
// Timestamp Field Tests
|
||||
// Schema: created_ts, updated_ts (timestamp, all comparison operators)
|
||||
// Functions: now(), arithmetic (+, -, *)
|
||||
// Time helpers: now variable, timestamp(...), duration(...)
|
||||
// =============================================================================
|
||||
|
||||
func TestMemoFilterCreatedTsComparison(t *testing.T) {
|
||||
|
|
@ -637,15 +637,15 @@ func TestMemoFilterCreatedTsComparison(t *testing.T) {
|
|||
tc.CreateMemo(NewMemoBuilder("memo-ts", tc.User.ID).Content("Timestamp test"))
|
||||
|
||||
// Test: created_ts < future (should match)
|
||||
memos := tc.ListWithFilter(`created_ts < ` + formatInt64(now+3600))
|
||||
memos := tc.ListWithFilter(`created_ts < timestamp(` + formatInt64(now+3600) + `)`)
|
||||
require.Len(t, memos, 1)
|
||||
|
||||
// Test: created_ts > past (should match)
|
||||
memos = tc.ListWithFilter(`created_ts > ` + formatInt64(now-3600))
|
||||
memos = tc.ListWithFilter(`created_ts > timestamp(` + formatInt64(now-3600) + `)`)
|
||||
require.Len(t, memos, 1)
|
||||
|
||||
// Test: created_ts > future (should not match)
|
||||
memos = tc.ListWithFilter(`created_ts > ` + formatInt64(now+3600))
|
||||
memos = tc.ListWithFilter(`created_ts > timestamp(` + formatInt64(now+3600) + `)`)
|
||||
require.Len(t, memos, 0)
|
||||
}
|
||||
|
||||
|
|
@ -656,12 +656,12 @@ func TestMemoFilterCreatedTsWithNow(t *testing.T) {
|
|||
|
||||
tc.CreateMemo(NewMemoBuilder("memo-ts-test", tc.User.ID).Content("Timestamp test"))
|
||||
|
||||
// Test: created_ts < now() + 5 (buffer for container clock drift)
|
||||
memos := tc.ListWithFilter(`created_ts < now() + 5`)
|
||||
// Test: created_ts < now + 5s (buffer for container clock drift)
|
||||
memos := tc.ListWithFilter(`created_ts < now + duration("5s")`)
|
||||
require.Len(t, memos, 1)
|
||||
|
||||
// Test: created_ts > now() + 5 (should not match)
|
||||
memos = tc.ListWithFilter(`created_ts > now() + 5`)
|
||||
// Test: created_ts > now + 5s (should not match)
|
||||
memos = tc.ListWithFilter(`created_ts > now + duration("5s")`)
|
||||
require.Len(t, memos, 0)
|
||||
}
|
||||
|
||||
|
|
@ -672,16 +672,16 @@ func TestMemoFilterCreatedTsArithmetic(t *testing.T) {
|
|||
|
||||
tc.CreateMemo(NewMemoBuilder("memo-ts-arith", tc.User.ID).Content("Timestamp arithmetic test"))
|
||||
|
||||
// Test: created_ts >= now() - 3600 (memos created in last hour)
|
||||
memos := tc.ListWithFilter(`created_ts >= now() - 3600`)
|
||||
// Test: created_ts >= now - 1h (memos created in last hour)
|
||||
memos := tc.ListWithFilter(`created_ts >= now - duration("1h")`)
|
||||
require.Len(t, memos, 1)
|
||||
|
||||
// Test: created_ts < now() - 86400 (memos older than 1 day - should be empty)
|
||||
memos = tc.ListWithFilter(`created_ts < now() - 86400`)
|
||||
// Test: created_ts < now - 24h (memos older than 1 day - should be empty)
|
||||
memos = tc.ListWithFilter(`created_ts < now - duration("24h")`)
|
||||
require.Len(t, memos, 0)
|
||||
|
||||
// Test: Multiplication - created_ts >= now() - 60 * 60
|
||||
memos = tc.ListWithFilter(`created_ts >= now() - 60 * 60`)
|
||||
// Test: chained duration arithmetic - created_ts >= now - 30m - 30m
|
||||
memos = tc.ListWithFilter(`created_ts >= now - duration("30m") - duration("30m")`)
|
||||
require.Len(t, memos, 1)
|
||||
}
|
||||
|
||||
|
|
@ -700,12 +700,12 @@ func TestMemoFilterUpdatedTs(t *testing.T) {
|
|||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Test: updated_ts >= now() - 60 (updated in last minute)
|
||||
memos := tc.ListWithFilter(`updated_ts >= now() - 60`)
|
||||
// Test: updated_ts >= now - 60s (updated in last minute)
|
||||
memos := tc.ListWithFilter(`updated_ts >= now - duration("60s")`)
|
||||
require.Len(t, memos, 1)
|
||||
|
||||
// Test: updated_ts > now() + 3600 (should be empty)
|
||||
memos = tc.ListWithFilter(`updated_ts > now() + 3600`)
|
||||
// Test: updated_ts > now + 1h (should be empty)
|
||||
memos = tc.ListWithFilter(`updated_ts > now + duration("1h")`)
|
||||
require.Len(t, memos, 0)
|
||||
}
|
||||
|
||||
|
|
@ -717,19 +717,19 @@ func TestMemoFilterAllComparisonOperators(t *testing.T) {
|
|||
tc.CreateMemo(NewMemoBuilder("memo-ops", tc.User.ID).Content("Comparison operators test"))
|
||||
|
||||
// Test: < (less than)
|
||||
memos := tc.ListWithFilter(`created_ts < now() + 3600`)
|
||||
memos := tc.ListWithFilter(`created_ts < now + duration("1h")`)
|
||||
require.Len(t, memos, 1)
|
||||
|
||||
// Test: <= (less than or equal) with buffer for clock drift
|
||||
memos = tc.ListWithFilter(`created_ts < now() + 5`)
|
||||
memos = tc.ListWithFilter(`created_ts < now + duration("5s")`)
|
||||
require.Len(t, memos, 1)
|
||||
|
||||
// Test: > (greater than)
|
||||
memos = tc.ListWithFilter(`created_ts > now() - 3600`)
|
||||
memos = tc.ListWithFilter(`created_ts > now - duration("1h")`)
|
||||
require.Len(t, memos, 1)
|
||||
|
||||
// Test: >= (greater than or equal)
|
||||
memos = tc.ListWithFilter(`created_ts >= now() - 60`)
|
||||
memos = tc.ListWithFilter(`created_ts >= now - duration("60s")`)
|
||||
require.Len(t, memos, 1)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -79,9 +79,10 @@ export const useMemoFilters = (options: UseMemoFiltersOptions = {}): string | un
|
|||
} else if (filter.factor === "displayTime") {
|
||||
const filterDate = new Date(filter.value);
|
||||
const filterUtcTimestamp = filterDate.getTime() + filterDate.getTimezoneOffset() * 60 * 1000;
|
||||
const timestampAfter = filterUtcTimestamp / 1000;
|
||||
const startTimestamp = Math.floor(filterUtcTimestamp / 1000);
|
||||
const endTimestamp = startTimestamp + 60 * 60 * 24;
|
||||
|
||||
conditions.push(`created_ts >= ${timestampAfter} && created_ts < ${timestampAfter + 60 * 60 * 24}`);
|
||||
conditions.push(`created_ts >= timestamp(${startTimestamp}) && created_ts < timestamp(${endTimestamp})`);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ const shortcutExamples = [
|
|||
},
|
||||
{
|
||||
title: "Recent notes",
|
||||
filter: "created_ts >= now() - 60 * 60",
|
||||
filter: 'created_ts >= now - duration("1h")',
|
||||
description: "Memos created in the last hour.",
|
||||
icon: Clock3Icon,
|
||||
},
|
||||
|
|
@ -103,6 +103,42 @@ const shortcutExamples = [
|
|||
description: "Every tag must satisfy the predicate (tagged memos only).",
|
||||
icon: TagsIcon,
|
||||
},
|
||||
{
|
||||
title: "One project tag",
|
||||
filter: 'tags.exists_one(t, t.startsWith("project/"))',
|
||||
description: "Exactly one tag matches the predicate.",
|
||||
icon: TagsIcon,
|
||||
},
|
||||
{
|
||||
title: "Any of these tags",
|
||||
filter: 'sets.intersects(tags, ["work", "urgent"])',
|
||||
description: "Tags intersect the given set.",
|
||||
icon: TagsIcon,
|
||||
},
|
||||
{
|
||||
title: "Exactly these tags",
|
||||
filter: 'sets.equivalent(tags, ["inbox"])',
|
||||
description: "Tagged with exactly this set, nothing more.",
|
||||
icon: TagsIcon,
|
||||
},
|
||||
{
|
||||
title: "Memos from 2024",
|
||||
filter: "created_ts.getFullYear() == 2024",
|
||||
description: "Filter by calendar year.",
|
||||
icon: Clock3Icon,
|
||||
},
|
||||
{
|
||||
title: "Weekend notes",
|
||||
filter: "created_ts.getDayOfWeek() == 0 || created_ts.getDayOfWeek() == 6",
|
||||
description: "Created on a Sunday or Saturday (0 = Sunday).",
|
||||
icon: Clock3Icon,
|
||||
},
|
||||
{
|
||||
title: "Long notes",
|
||||
filter: "size(content) > 280",
|
||||
description: "Memos longer than 280 characters.",
|
||||
icon: FilterIcon,
|
||||
},
|
||||
];
|
||||
|
||||
const filterFields = [
|
||||
|
|
@ -115,12 +151,22 @@ const filterFields = [
|
|||
"tag in [...]",
|
||||
"tags.exists(...)",
|
||||
"tags.all(...)",
|
||||
"tags.exists_one(...)",
|
||||
"sets.contains(tags, [...])",
|
||||
"sets.intersects(tags, [...])",
|
||||
"sets.equivalent(tags, [...])",
|
||||
"size(content) > ...",
|
||||
"has_task_list",
|
||||
"has_incomplete_tasks",
|
||||
"has_link",
|
||||
"has_code",
|
||||
"created_ts",
|
||||
'created_ts >= now - duration("24h")',
|
||||
"created_ts.getFullYear() == ...",
|
||||
"created_ts.getMonth() == ... (0 = Jan)",
|
||||
"created_ts.getDayOfWeek() == ... (0 = Sun)",
|
||||
"updated_ts",
|
||||
"now",
|
||||
'timestamp("2025-01-01T00:00:00Z")',
|
||||
];
|
||||
|
||||
const getShortcutId = (name: string): string => {
|
||||
|
|
@ -435,7 +481,11 @@ const Shortcuts = () => {
|
|||
/>
|
||||
<p className="text-xs leading-5 text-muted-foreground">
|
||||
Combine expressions with <span className="font-mono">&&</span>, <span className="font-mono">||</span>, and{" "}
|
||||
<span className="font-mono">!</span>. Time fields use Unix seconds and support <span className="font-mono">now()</span>.
|
||||
<span className="font-mono">!</span>. Time fields are timestamps — use <span className="font-mono">now</span>,{" "}
|
||||
<span className="font-mono">duration("24h")</span>, <span className="font-mono">timestamp(...)</span>, and accessors
|
||||
like <span className="font-mono">created_ts.getFullYear()</span>. Tags support{" "}
|
||||
<span className="font-mono">sets.contains/intersects/equivalent</span> and{" "}
|
||||
<span className="font-mono">size(content)</span> measures length.
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
|
|
|
|||
Loading…
Reference in a new issue