feat: add AI matching, classification, and scoring functions (#9300)

* feat: return matching scores from jev

Replace the experimental three-argument Boolean function with jev(text, prompt) returning a Float64 probability in [0, 1]. Move threshold comparisons into SQL and update tests and migration examples.

Signed-off-by: Lei, HUANG <ratuthomm@gmail.com>

* feat: add Jev choice and score functions

Share asynchronous execution across Noul, Choice, and Score. Validate JSON criteria before requests and return typed scalar answers. Add SQL and HTTP mock coverage with usage examples.

Signed-off-by: Lei, HUANG <ratuthomm@gmail.com>

* refactor: use generic AI SQL function names

Expose ai_match, ai_choose, and ai_score and move their implementation, tests, and usage guide under generic AI names. Document the current unreleased interface without migration history.

Signed-off-by: Lei, HUANG <ratuthomm@gmail.com>

* fix: share constant AI criteria within each batch

Borrow scalar string arguments and lazily parse constant criteria once per batch. Share the parsed allocation across requests while preserving NULL propagation and batch validation before HTTP calls.

Signed-off-by: Lei, HUANG <ratuthomm@gmail.com>

* feat: preserve AI score uncertainty in JSONB results

Return score, confidence, and probabilities in criteria-level order from one evaluation. Validate the distribution and preserve provider precision. Add JSON extraction, uncertainty, and single-request regressions, and document confidence-aware ranking.

Signed-off-by: Lei, HUANG <ratuthomm@gmail.com>

* docs: explain reuse of volatile AI evaluations

Document repeated SELECT and WHERE evaluation costs as N + M requests, and show subquery aliases for reusing scalar or structured AI results without additional model calls.

Signed-off-by: Lei, HUANG <ratuthomm@gmail.com>

---------

Signed-off-by: Lei, HUANG <ratuthomm@gmail.com>
This commit is contained in:
Lei, HUANG
2026-09-22 14:12:24 +00:00
committed by GitHub
parent a9e2a89b7b
commit d0f8f4b80c
12 changed files with 1692 additions and 686 deletions
+247
View File
@@ -0,0 +1,247 @@
# AI functions (experimental MVP)
Three asynchronous SQL scalar functions provide natural-language matching,
classification, and rating. The current backend is TypeSafe's Jev model:
| Function | Mode | Result |
| --- | --- | --- |
| `ai_match(text, prompt)` | Noul | `Float64` probability in **[0, 1]** that the prompt's statement is true |
| `ai_choose(text, prompt, criteria)` | Choice | `String` containing the selected option's name |
| `ai_score(text, prompt, criteria)` | Score | JSONB object containing `score`, `confidence`, and a `probabilities` array |
Use scalar results or extracted JSON fields in `SELECT`, comparisons in `WHERE`,
and `ORDER BY`.
All arguments are SQL strings. Any SQL NULL argument produces SQL `NULL`
without validating that row's criteria or calling the API.
## Start GreptimeDB
The functions are compiled and registered by the **default-enabled Cargo feature
`ai_functions`**. Enable the Jev backend separately in the server process:
```sh
export GREPTIMEDB_EXPERIMENTAL_JEV=true
export JEV_API_KEY='<your TypeSafe API key>'
cargo run -p cmd -- standalone start
```
If your key is already exported in `~/.zshrc`, run `source ~/.zshrc` first.
The optional `JEV_MODEL` defaults to `jev-latest`; `JEV_ENDPOINT` defaults to
`https://api.typesafe.ai/v1/systemone` (the full evaluation endpoint URL).
The MVP uses environment variables rather than TOML configuration.
The Cargo feature includes all three SQL functions in the build; the runtime environment
variables enable and configure its API calls. Default compilation does not turn
on external API calls. To enable the feature explicitly, use
`--features ai_functions`.
## Query
For a new example table, quote `service` and `message` in the DDL because the
current parser treats them as keywords:
```sql
CREATE TABLE events (
occurred_at TIMESTAMP TIME INDEX,
"service" STRING,
"message" STRING,
PRIMARY KEY ("service")
);
```
### Matching probability (Noul)
```sql
SELECT occurred_at, service, message
FROM events
WHERE occurred_at >= '2026-09-19T00:00:00Z'
AND occurred_at < '2026-09-20T00:00:00Z'
AND service = 'payments'
AND ai_match(
(message),
'The event reports that a payment still failed after retries.'
) >= 0.8
ORDER BY occurred_at;
```
To return the scores and rank matching events:
```sql
SELECT occurred_at, message,
ai_match(message, 'The event reports that a payment still failed after retries.') AS score
FROM events
WHERE occurred_at >= '2026-09-19T00:00:00Z'
AND occurred_at < '2026-09-20T00:00:00Z'
AND service = 'payments'
ORDER BY score DESC NULLS LAST;
```
Both arguments are strings; larger scores indicate a stronger match.
`(message)` is an ordinary parenthesized string expression. For several columns,
combine them explicitly, for example `concat(service, ': ', message)`.
Any null argument produces SQL `NULL`, which `WHERE` excludes, without an API call.
### Classification (Choice)
`criteria` is a JSON object with **1 to 255 options**. Each key is an option name;
its value is a description (string, object, or array), or JSON `null` if the name
is sufficient. The function returns the selected key, suitable for filtering
or grouping:
```sql
SELECT occurred_at, message,
ai_choose(message, 'Which team should handle this event?',
'{"billing":"Payments, invoices, refunds","technical":"Bugs and outages","other":null}') AS team
FROM events
WHERE occurred_at >= '2026-09-19T00:00:00Z'
AND occurred_at < '2026-09-20T00:00:00Z';
```
### Rating (Score)
`criteria` is a JSON array of **2 to 10 ordered level descriptions**, from low
to high. Descriptions can be strings, objects, or arrays. Level numbers start
at zero. The function returns a JSONB object from one model evaluation:
```json
{
"score": 1.05,
"confidence": 0.92,
"probabilities": [0.0, 0.95, 0.05]
}
```
- `score` is the provider's probability-weighted mean of the level numbers, in
`[0, N − 1]` for N levels. It may fall between levels.
- `confidence` is the provider's confidence in `[0, 1]`, reflecting how concentrated
the distribution is. It is not a guarantee that the answer is correct.
- `probabilities[i]` is the probability of `criteria[i]`, with exactly N entries in
level order. These are individual level probabilities, not cumulative probabilities.
Different distributions can have the same score. `[0, 1, 0]` and `[0.5, 0, 0.5]`
both have score `1`: the first concentrates on the middle level, while the second
is split between the extremes. Inspect confidence and the distribution before
treating a mean score as a definite severity level.
For example, compute the rating once, extract numeric fields, and rank only rows
meeting a chosen confidence threshold:
```sql
SELECT occurred_at, message,
json_get_float(rating, 'score') AS severity,
json_get_float(rating, 'confidence') AS confidence
FROM (
SELECT occurred_at, message,
ai_score(message, 'How severe is this event?',
'["No impact to functionality","Degraded service with a workaround","Blocking issue with no workaround"]') AS rating
FROM events
WHERE occurred_at >= '2026-09-19T00:00:00Z'
AND occurred_at < '2026-09-20T00:00:00Z'
) AS rated
WHERE json_get_float(rating, 'confidence') >= 0.8
ORDER BY severity DESC NULLS LAST;
```
Use `json_to_string(rating)` to display the full object, or
`json_get_float(rating, 'probabilities[2]')` to read the probability of level 2.
The array positions correspond to the supplied criteria, which provide the level
descriptions.
## Reusing an evaluation
AI functions are marked volatile to prevent HTTP calls during planning. Identical
calls written separately in `SELECT` and `WHERE` are evaluated separately. If the
filter evaluates N non-null rows and M rows pass to the projection, this can cost
N + M requests, reaching 2N when all rows pass.
To return a result and filter by it, compute it once in a subquery and reference
its alias in the outer query:
```sql
SELECT occurred_at, message, score
FROM (
SELECT occurred_at, message,
ai_match(message, 'The event reports that a payment still failed after retries.') AS score
FROM events
WHERE occurred_at >= '2026-09-19T00:00:00Z'
AND occurred_at < '2026-09-20T00:00:00Z'
AND service = 'payments'
) AS scored
WHERE score >= 0.8
ORDER BY score DESC;
```
This also applies to `ai_choose` and `ai_score`. Extract multiple JSON fields from
the aliased `ai_score` result, as in the rating example above, rather than calling
the model again for each field. Filter pushdown preserves the volatile projection
instead of duplicating it into the outer predicate.
## MVP behavior
- Each non-null row makes one HTTP request per function invocation: the text is
`state`, and the prompt is the question's `instructions`. Choice and Score pass
the parsed JSON as `criteria`. Calls to different functions are separate requests.
- Criteria for all non-null rows in a batch are validated before sending that
batch's requests. Malformed JSON, invalid description types, or invalid option/
level counts fail the query locally.
- Answers must have the requested question type. Noul probabilities outside
`[0, 1]`, Choice labels not in the criteria, and Score values outside `[0, N − 1]`
fail the query rather than being clamped or replaced with NULL. Score responses
must also contain a numeric confidence in `[0, 1]` and probabilities for every
level. Each probability must be numeric and in `[0, 1]`; their sum must be within
`1e-6` of 1. Provider values are preserved without renormalizing the distribution
or recomputing the score.
- Up to eight requests run concurrently per expression/batch invocation, not per
query or process. Concurrent partitions and queries can exceed eight requests
in total. Each request has a 30-second timeout. Errors (including rate limits
and invalid responses) fail the query. There is no automatic retry or
cross-query cache.
- Start with small, time-bounded queries. Ordinary SQL predicates can reduce the
candidate set, but SQL does not guarantee left-to-right predicate evaluation or
that `LIMIT` bounds the number of API calls.
- Text passed to the function is sent to the configured TypeSafe endpoint.
The API key remains in the server environment, not in SQL.
- This MVP is intended for standalone use. A distributed deployment would need
the environment on every node evaluating the function and separate validation.
API reference: <https://docs.typesafe.ai/api>
## Before stabilizing
The experimental MVP still needs the following controls before stabilization:
- A process-wide concurrency limit shared by AI function invocations, such as a semaphore.
- Bounded retries with exponential backoff for HTTP `429` and `529`, following
TypeSafe's rate-limit guidance.
- Request budgets and metrics for API calls, latency, retries, and rate-limit errors.
These are follow-up work, not guarantees provided by the current implementation.
## Validation
Check both the default build (where all three functions are registered) and an isolated
`common-function` build without default features. The regular AI function tests use a local
HTTP server and need no API key:
```sh
cargo nextest run -p common-function
cargo nextest run -p common-function --no-default-features
cargo nextest run -p common-function --no-default-features --features ai_functions
```
The sqlness runner explicitly includes `ai_functions` when building its test binary:
```sh
cargo sqlness bare -t ai_functions
```
If using `--bins-dir`, provide a binary built with `ai_functions` (included by
default). CI covers the AI-enabled path through the regular unit tests and
sqlness tests; the feature-off checks above can be run locally when needed.
An opt-in test calls the real service on synthetic payment events:
```sh
source ~/.zshrc
GREPTIMEDB_EXPERIMENTAL_JEV=true cargo nextest run -p common-function \
-E 'test(test_ai_match_live)' --run-ignored ignored-only
```
-119
View File
@@ -1,119 +0,0 @@
# Natural-language filtering with Jev (experimental MVP)
`jev(text, statement, threshold)` returns whether Jev judges the statement true
for the text with probability **greater than or equal to** the threshold.
It is an asynchronous SQL scalar function usable in `WHERE` and `SELECT`.
## Start GreptimeDB
Jev is compiled and registered by the **default-enabled Cargo feature
`ai_functions`**. Enable API evaluation separately in the server process:
```sh
export GREPTIMEDB_EXPERIMENTAL_JEV=true
export JEV_API_KEY='<your TypeSafe API key>'
cargo run -p cmd -- standalone start
```
If your key is already exported in `~/.zshrc`, run `source ~/.zshrc` first.
The optional `JEV_MODEL` defaults to `jev-latest`; `JEV_ENDPOINT` defaults to
`https://api.typesafe.ai/v1/systemone` (the full evaluation endpoint URL).
The MVP uses environment variables rather than TOML configuration.
The Cargo feature includes the SQL function in the build; the runtime environment
variables enable and configure its API calls. Default compilation does not turn
on external API calls. To enable the feature explicitly, use
`--features ai_functions`.
## Query
For a new example table, quote `service` and `message` in the DDL because the
current parser treats them as keywords:
```sql
CREATE TABLE events (
occurred_at TIMESTAMP TIME INDEX,
"service" STRING,
"message" STRING,
PRIMARY KEY ("service")
);
```
```sql
SELECT occurred_at, service, message
FROM events
WHERE occurred_at >= '2026-09-19T00:00:00Z'
AND occurred_at < '2026-09-20T00:00:00Z'
AND service = 'payments'
AND jev(
(message),
'The event reports that a payment still failed after retries.',
0.8
)
ORDER BY occurred_at;
```
The first two arguments are strings and the third is a number in `[0, 1]`.
`(message)` is an ordinary parenthesized string expression. For several columns,
combine them explicitly, for example `concat(service, ': ', message)`.
Any null argument produces SQL `NULL`, which `WHERE` excludes, without an API call.
## MVP behavior
- Each non-null row makes one HTTP request: the text is `state`, and the statement
is a `noul` question's `instructions`. The returned `noul` probability is compared
with the threshold locally; this is not the separate Choice/Score confidence.
- Up to eight requests run concurrently per expression/batch invocation, not per
query or process. Concurrent partitions and queries can exceed eight requests
in total. Each request has a 30-second timeout. Errors (including rate limits
and invalid responses) fail the query. There is no automatic retry or
cross-query cache.
- Start with small, time-bounded queries. Ordinary SQL predicates can reduce the
candidate set, but SQL does not guarantee left-to-right predicate evaluation or
that `LIMIT` bounds the number of API calls.
- Text passed to the function is sent to the configured TypeSafe endpoint.
The API key remains in the server environment, not in SQL.
- This MVP is intended for standalone use. A distributed deployment would need
the environment on every node evaluating the function and separate validation.
API reference: <https://docs.typesafe.ai/api>
## Before stabilizing
The experimental MVP still needs the following controls before stabilization:
- A process-wide concurrency limit shared by Jev invocations, such as a semaphore.
- Bounded retries with exponential backoff for HTTP `429` and `529`, following
TypeSafe's rate-limit guidance.
- Request budgets and metrics for API calls, latency, retries, and rate-limit errors.
These are follow-up work, not guarantees provided by the current implementation.
## Validation
Check both the default build (where `jev` is registered) and an isolated
`common-function` build without default features. Jev's regular tests use a local
HTTP server and need no API key:
```sh
cargo nextest run -p common-function
cargo nextest run -p common-function --no-default-features
cargo nextest run -p common-function --no-default-features --features ai_functions
```
The sqlness runner explicitly includes `ai_functions` when building its test binary:
```sh
cargo sqlness bare -t jev
```
If using `--bins-dir`, provide a binary built with `ai_functions` (included by
default). CI covers the AI-enabled path through the regular unit tests and
sqlness tests; the feature-off checks above can be run locally when needed.
An opt-in test calls the real service on synthetic payment events:
```sh
source ~/.zshrc
GREPTIMEDB_EXPERIMENTAL_JEV=true cargo nextest run -p common-function \
-E 'test(test_jev_live)' --run-ignored ignored-only
```
+14 -8
View File
@@ -28,14 +28,14 @@ use crate::aggrs::count_hash::CountHash;
use crate::aggrs::vector::VectorFunction as VectorAggrFunction;
use crate::function::{Function, FunctionRef};
use crate::function_factory::ScalarFunctionFactory;
#[cfg(feature = "ai_functions")]
use crate::scalars::ai;
use crate::scalars::anomaly::AnomalyFunction;
use crate::scalars::avg_calc::AvgCalcFunction;
use crate::scalars::date::DateFunction;
use crate::scalars::expression::ExpressionFunction;
use crate::scalars::hll_count::HllCalcFunction;
use crate::scalars::ip::IpFunctions;
#[cfg(feature = "ai_functions")]
use crate::scalars::jev::JevFunction;
use crate::scalars::json::JsonFunction;
use crate::scalars::matches::MatchesFunction;
use crate::scalars::matches_term::MatchesTermFunction;
@@ -229,7 +229,7 @@ pub static FUNCTION_REGISTRY: LazyLock<Arc<FunctionRegistry>> = LazyLock::new(||
MatchesFunction::register(&function_registry);
MatchesTermFunction::register(&function_registry);
#[cfg(feature = "ai_functions")]
JevFunction::register(&function_registry);
ai::register(&function_registry);
// System and administration functions
SystemFunction::register(&function_registry);
@@ -384,11 +384,17 @@ mod tests {
}
#[test]
fn test_jev_registration_matches_feature() {
assert_eq!(
FUNCTION_REGISTRY.get_function("jev").is_some(),
cfg!(feature = "ai_functions")
);
fn test_ai_registration_matches_feature() {
for name in ["ai_match", "ai_choose", "ai_score"] {
assert_eq!(
FUNCTION_REGISTRY.get_function(name).is_some(),
cfg!(feature = "ai_functions"),
"{name}"
);
}
for name in ["jev", "jev_choice", "jev_score"] {
assert!(FUNCTION_REGISTRY.get_function(name).is_none(), "{name}");
}
}
#[test]
+2 -2
View File
@@ -12,13 +12,13 @@
// See the License for the specific language governing permissions and
// limitations under the License.
#[cfg(feature = "ai_functions")]
pub(crate) mod ai;
pub mod anomaly;
pub(crate) mod date;
pub mod expression;
#[cfg(feature = "geo")]
pub mod geo;
#[cfg(feature = "ai_functions")]
pub(crate) mod jev;
pub mod json;
pub mod matches;
pub mod matches_term;
+454
View File
@@ -0,0 +1,454 @@
// Copyright 2023 Greptime Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#[cfg(test)]
mod tests;
use std::fmt;
use std::hash::Hash;
use std::marker::PhantomData;
use std::sync::Arc;
use std::time::Duration;
use arrow::array::{Array, StringViewArray, new_null_array};
use arrow::datatypes::DataType;
use async_trait::async_trait;
use datafusion_common::cast::as_string_view_array;
use datafusion_common::{Result, ScalarValue, exec_datafusion_err, exec_err, not_impl_err};
use datafusion_expr::async_udf::{AsyncScalarUDF, AsyncScalarUDFImpl};
use datafusion_expr::{ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl, Signature, Volatility};
use futures::{StreamExt, TryStreamExt, stream};
use reqwest::Client;
use serde::Serialize;
use serde_json::{Map, Value, json};
use crate::function_factory::ScalarFunctionFactory;
use crate::function_registry::FunctionRegistry;
/// Registers the experimental AI matching, classification, and rating functions.
pub(crate) fn register(registry: &FunctionRegistry) {
AiFunction::<Noul>::register(registry);
AiFunction::<Choice>::register(registry);
AiFunction::<Score>::register(registry);
}
/// Shared asynchronous execution for AI functions using the Jev backend.
#[derive(PartialEq, Eq, Hash)]
struct AiFunction<Q> {
signature: Signature,
question: PhantomData<Q>,
enabled: bool,
api_key: Option<String>,
endpoint: String,
model: String,
}
struct AiRequest<'a, Q: AiQuestion> {
text: &'a str,
prompt: &'a str,
criteria: Arc<Q::Criteria>,
}
/// Borrows string arguments without expanding constants to the batch size.
enum StringArgument<'a> {
Scalar(Option<&'a str>),
Array(&'a StringViewArray),
}
impl<'a> StringArgument<'a> {
fn try_new(arg: &'a ColumnarValue) -> Result<Self> {
match arg {
ColumnarValue::Scalar(ScalarValue::Utf8View(value)) => {
Ok(Self::Scalar(value.as_deref()))
}
ColumnarValue::Array(array) => Ok(Self::Array(as_string_view_array(array)?)),
_ => exec_err!("AI functions expect Utf8View arguments"),
}
}
fn value(&self, row: usize) -> Option<&'a str> {
match self {
Self::Scalar(value) => *value,
Self::Array(array) => {
let array = *array;
array.is_valid(row).then(|| array.value(row))
}
}
}
}
impl<Q: AiQuestion> AiFunction<Q> {
fn register(registry: &FunctionRegistry) {
registry.register(ScalarFunctionFactory {
name: Q::NAME.to_string(),
factory: Arc::new(|_| AsyncScalarUDF::new(Arc::new(Self::default())).into_scalar_udf()),
});
}
async fn evaluate(&self, client: &Client, request: AiRequest<'_, Q>) -> Result<ScalarValue> {
let mut question = json!({ "type": Q::TYPE, "instructions": request.prompt });
if Q::ARG_COUNT == 3 {
question["criteria"] = json!(request.criteria.as_ref());
}
let response: Value = client
.post(&self.endpoint)
.json(&json!({
"model": self.model,
"state": request.text,
"questions": { "matches": question }
}))
.send()
.await
.and_then(|response| response.error_for_status())
.map_err(|e| exec_datafusion_err!("{} request failed: {e}", Q::NAME))?
.json()
.await
.map_err(|e| exec_datafusion_err!("{} response is not valid JSON: {e}", Q::NAME))?;
let answer = &response["answers"]["matches"];
if answer["type"].as_str() != Some(Q::TYPE) {
return exec_err!(
"{} response must contain answers.matches with type {}",
Q::NAME,
Q::TYPE
);
}
Q::parse_answer(answer, request.criteria.as_ref())
}
fn prepare_requests<'a>(
&self,
args: &'a [ColumnarValue],
number_rows: usize,
) -> Result<Vec<Option<AiRequest<'a, Q>>>> {
if args.len() != Q::ARG_COUNT {
return exec_err!("{} requires {} arguments", Q::NAME, Q::ARG_COUNT);
}
let args = args
.iter()
.map(StringArgument::try_new)
.collect::<Result<Vec<_>>>()?;
let scalar_criteria = args[2..]
.iter()
.all(|arg| matches!(arg, StringArgument::Scalar(_)));
let mut shared_criteria: Option<Arc<Q::Criteria>> = None;
// Validate all non-null rows before making any billable requests.
(0..number_rows)
.map(|row| {
let values: Option<Vec<_>> = args.iter().map(|arg| arg.value(row)).collect();
values
.map(|values| {
// Initialize lazily so NULL rows never validate otherwise invalid criteria.
let criteria = match &shared_criteria {
Some(criteria) => Arc::clone(criteria),
None => {
let criteria = Arc::new(Q::parse_criteria(&values[2..])?);
if scalar_criteria {
shared_criteria = Some(Arc::clone(&criteria));
}
criteria
}
};
Ok(AiRequest {
text: values[0],
prompt: values[1],
criteria,
})
})
.transpose()
})
.collect()
}
fn client(&self) -> Result<Client> {
if !self.enabled {
return exec_err!(
"{} is experimental; set GREPTIMEDB_EXPERIMENTAL_JEV=true to enable it",
Q::NAME
);
}
let key = self
.api_key
.as_deref()
.filter(|key| !key.trim().is_empty())
.ok_or_else(|| {
exec_datafusion_err!("{} requires the JEV_API_KEY environment variable", Q::NAME)
})?;
let mut authorization = reqwest::header::HeaderValue::from_str(&format!("Bearer {key}"))
.map_err(|_| exec_datafusion_err!("JEV_API_KEY is not a valid HTTP header value"))?;
authorization.set_sensitive(true);
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(reqwest::header::AUTHORIZATION, authorization);
Client::builder()
.default_headers(headers)
.timeout(Duration::from_secs(30))
.build()
.map_err(|e| exec_datafusion_err!("failed to create {} HTTP client: {e}", Q::NAME))
}
}
impl<Q: AiQuestion> Default for AiFunction<Q> {
fn default() -> Self {
Self {
// External model evaluations must not be constant-folded during planning.
signature: Signature::exact(
vec![DataType::Utf8View; Q::ARG_COUNT],
Volatility::Volatile,
),
question: PhantomData,
enabled: std::env::var("GREPTIMEDB_EXPERIMENTAL_JEV").as_deref() == Ok("true"),
api_key: std::env::var("JEV_API_KEY").ok(),
endpoint: std::env::var("JEV_ENDPOINT")
.unwrap_or_else(|_| "https://api.typesafe.ai/v1/systemone".to_string()),
model: std::env::var("JEV_MODEL").unwrap_or_else(|_| "jev-latest".to_string()),
}
}
}
impl<Q: AiQuestion> fmt::Debug for AiFunction<Q> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
// Credentials must never appear in plans or diagnostic output.
f.debug_struct("AiFunction")
.field("name", &Q::NAME)
.field("signature", &self.signature)
.field("enabled", &self.enabled)
.field("model", &self.model)
.finish_non_exhaustive()
}
}
impl<Q: AiQuestion> ScalarUDFImpl for AiFunction<Q> {
fn name(&self) -> &str {
Q::NAME
}
fn signature(&self) -> &Signature {
&self.signature
}
fn return_type(&self, _arg_types: &[DataType]) -> Result<DataType> {
Ok(Q::return_type())
}
fn invoke_with_args(&self, _args: ScalarFunctionArgs) -> Result<ColumnarValue> {
not_impl_err!("{} can only be called from async contexts", Q::NAME)
}
}
#[async_trait]
impl<Q: AiQuestion> AsyncScalarUDFImpl for AiFunction<Q> {
async fn invoke_async_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
let rows = self.prepare_requests(&args.args, args.number_rows)?;
if rows.iter().all(Option::is_none) {
return Ok(ColumnarValue::Array(new_null_array(
&Q::return_type(),
args.number_rows,
)));
}
// Reuse connections within the batch; bound in-flight requests and preserve row order.
let client = self.client()?;
// Collect futures to avoid async_trait's higher-ranked lifetime inference issue.
let requests: Vec<_> = rows
.into_iter()
.map(|row| {
let client = &client;
async move {
match row {
Some(request) => self.evaluate(client, request).await,
None => ScalarValue::try_from(&Q::return_type()),
}
}
})
.collect();
// This limit is per expression/batch invocation, not per query or process.
// Concurrent partitions and queries can each have their own in-flight requests.
let answers: Vec<ScalarValue> = stream::iter(requests).buffered(8).try_collect().await?;
Ok(ColumnarValue::Array(ScalarValue::iter_to_array(answers)?))
}
}
/// The request criteria and scalar answer contract for an AI question type.
trait AiQuestion: fmt::Debug + Eq + Hash + Send + Sync + 'static {
type Criteria: Serialize + Send + Sync;
const NAME: &'static str;
const TYPE: &'static str;
const ARG_COUNT: usize;
fn return_type() -> DataType;
fn parse_criteria(args: &[&str]) -> Result<Self::Criteria>;
fn parse_answer(answer: &Value, criteria: &Self::Criteria) -> Result<ScalarValue>;
}
#[derive(Debug, PartialEq, Eq, Hash)]
struct Noul;
impl AiQuestion for Noul {
type Criteria = ();
const NAME: &'static str = "ai_match";
const TYPE: &'static str = "noul";
const ARG_COUNT: usize = 2;
fn return_type() -> DataType {
DataType::Float64
}
fn parse_criteria(_args: &[&str]) -> Result<()> {
Ok(())
}
fn parse_answer(answer: &Value, _criteria: &()) -> Result<ScalarValue> {
numeric_answer(Self::NAME, &answer["noul"], "noul", 1.0)
.map(|probability| ScalarValue::Float64(Some(probability)))
}
}
#[derive(Debug, PartialEq, Eq, Hash)]
struct Choice;
impl AiQuestion for Choice {
type Criteria = Map<String, Value>;
const NAME: &'static str = "ai_choose";
const TYPE: &'static str = "choice";
const ARG_COUNT: usize = 3;
fn return_type() -> DataType {
DataType::Utf8
}
fn parse_criteria(args: &[&str]) -> Result<Self::Criteria> {
let criteria: Value = serde_json::from_str(args[0])
.map_err(|e| exec_datafusion_err!("ai_choose criteria is not valid JSON: {e}"))?;
match criteria {
Value::Object(options)
if (1..=255).contains(&options.len())
&& options.values().all(|v| v.is_null() || is_description(v)) =>
{
Ok(options)
}
_ => exec_err!(
"ai_choose criteria must be a JSON object with 1 to 255 options; descriptions must be strings, objects, arrays, or null"
),
}
}
fn parse_answer(answer: &Value, criteria: &Self::Criteria) -> Result<ScalarValue> {
let choice = answer["choice"]
.as_str()
.filter(|choice| criteria.contains_key(*choice))
.ok_or_else(|| {
exec_datafusion_err!("ai_choose response must contain a choice from the criteria")
})?;
Ok(ScalarValue::Utf8(Some(choice.to_string())))
}
}
#[derive(Debug, PartialEq, Eq, Hash)]
struct Score;
impl AiQuestion for Score {
type Criteria = Vec<Value>;
const NAME: &'static str = "ai_score";
const TYPE: &'static str = "score";
const ARG_COUNT: usize = 3;
fn return_type() -> DataType {
DataType::BinaryView
}
fn parse_criteria(args: &[&str]) -> Result<Self::Criteria> {
let criteria: Value = serde_json::from_str(args[0])
.map_err(|e| exec_datafusion_err!("ai_score criteria is not valid JSON: {e}"))?;
match criteria {
Value::Array(levels)
if (2..=10).contains(&levels.len()) && levels.iter().all(is_description) =>
{
Ok(levels)
}
_ => exec_err!(
"ai_score criteria must be a JSON array with 2 to 10 levels; descriptions must be strings, objects, or arrays"
),
}
}
fn parse_answer(answer: &Value, criteria: &Self::Criteria) -> Result<ScalarValue> {
let score = numeric_answer(
Self::NAME,
&answer["score"],
"score",
(criteria.len() - 1) as f64,
)?;
let confidence = numeric_answer(Self::NAME, &answer["confidence"], "confidence", 1.0)?;
let probabilities = score_probabilities(&answer["probabilities"], criteria.len())?;
let object = jsonb::Object::from([
("score".to_string(), jsonb::Value::from(score)),
("confidence".to_string(), jsonb::Value::from(confidence)),
(
"probabilities".to_string(),
jsonb::Value::Array(probabilities.into_iter().map(jsonb::Value::from).collect()),
),
]);
Ok(ScalarValue::BinaryView(Some(
jsonb::Value::Object(object).to_vec(),
)))
}
}
fn score_probabilities(distribution: &Value, level_count: usize) -> Result<Vec<f64>> {
if distribution
.as_object()
.is_none_or(|probabilities| probabilities.len() != level_count)
{
return exec_err!(
"ai_score response must contain probabilities for all {level_count} levels"
);
}
let probabilities = (0..level_count)
.map(|level| {
numeric_answer(
Score::NAME,
&distribution[level.to_string()],
"probability",
1.0,
)
})
.collect::<Result<Vec<_>>>()?;
// Allow small rounding differences without renormalizing the provider's distribution.
if (probabilities.iter().sum::<f64>() - 1.0).abs() > 1e-6 {
return exec_err!("ai_score response probabilities must sum to 1 within 1e-6");
}
Ok(probabilities)
}
fn is_description(description: &Value) -> bool {
matches!(
description,
Value::String(_) | Value::Object(_) | Value::Array(_)
)
}
fn numeric_answer(name: &str, answer: &Value, field: &str, max: f64) -> Result<f64> {
answer
.as_f64()
.filter(|score| (0.0..=max).contains(score))
.ok_or_else(|| {
exec_datafusion_err!("{name} response must contain a finite {field} in [0, {max}]")
})
}
+766
View File
@@ -0,0 +1,766 @@
// Copyright 2023 Greptime Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::sync::atomic::{AtomicUsize, Ordering};
use axum::http::{HeaderMap, StatusCode};
use axum::routing::post;
use axum::{Json, Router};
use datafusion::assert_batches_eq;
use datafusion::prelude::{SessionConfig, SessionContext};
use tokio::net::TcpListener;
use tokio::task::JoinHandle;
use super::*;
use crate::function::FunctionContext;
use crate::function_registry::FUNCTION_REGISTRY;
struct MockServer {
endpoint: String,
task: JoinHandle<()>,
}
impl MockServer {
async fn start(router: Router) -> Self {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let endpoint = format!("http://{}/v1/systemone", listener.local_addr().unwrap());
let task = tokio::spawn(async move { axum::serve(listener, router).await.unwrap() });
Self { endpoint, task }
}
fn function<Q: AiQuestion>(&self) -> AiFunction<Q> {
AiFunction {
enabled: true,
api_key: Some("test-key".to_string()),
endpoint: self.endpoint.clone(),
model: "test-model".to_string(),
..Default::default()
}
}
}
impl Drop for MockServer {
fn drop(&mut self) {
self.task.abort();
}
}
fn context<Q: AiQuestion>(function: AiFunction<Q>) -> SessionContext {
let ctx = SessionContext::new_with_config(SessionConfig::new().with_target_partitions(1));
ctx.register_udf(AsyncScalarUDF::new(Arc::new(function)).into_scalar_udf());
let json_get_float = FUNCTION_REGISTRY.get_function("json_get_float").unwrap();
ctx.register_udf(json_get_float.provide(FunctionContext::default()));
ctx
}
#[test]
fn test_ai_scalar_criteria_share_allocation_and_skip_null_rows() {
fn assert_shared_criteria<Q: AiQuestion>(criteria: &str) {
let function = AiFunction::<Q>::default();
for (texts, criteria) in [
(vec![None, Some("first"), Some("second")], criteria),
(vec![None, None], "not json"),
(vec![], "not json"),
] {
let number_rows = texts.len();
let expected_requests = texts.iter().flatten().count();
let args = vec![
ColumnarValue::Array(Arc::new(StringViewArray::from(texts))),
ColumnarValue::Scalar(ScalarValue::Utf8View(Some("prompt".to_string()))),
ColumnarValue::Scalar(ScalarValue::Utf8View(Some(criteria.to_string()))),
];
let requests = function.prepare_requests(&args, number_rows).unwrap();
assert_eq!(requests.len(), number_rows);
assert_eq!(requests.iter().flatten().count(), expected_requests);
// Regression: parsed constant criteria must not be duplicated for each row.
let mut criteria = requests.iter().flatten().map(|request| &request.criteria);
if let Some(first) = criteria.next() {
assert!(criteria.all(|other| Arc::ptr_eq(first, other)));
}
}
}
assert_shared_criteria::<Choice>(r#"{"billing":null,"technical":"errors"}"#);
assert_shared_criteria::<Score>(r#"["low","high"]"#);
}
#[tokio::test]
async fn test_ai_match_sql_returns_and_filters_noul_probability() {
let server = MockServer::start(Router::new().route(
"/v1/systemone",
post(
|headers: HeaderMap, Json(request): Json<Value>| async move {
assert_eq!(headers["authorization"], "Bearer test-key");
assert_eq!(request["model"], "test-model");
assert_eq!(request["questions"]["matches"]["type"], "noul");
// The prompt is passed through verbatim.
assert_eq!(
request["questions"]["matches"]["instructions"],
"payment failed"
);
let probability = match request["state"].as_str().unwrap() {
"failed" => {
// Complete this row after later rows to exercise result alignment.
tokio::time::sleep(Duration::from_millis(20)).await;
0.95
}
"boundary" => 0.8,
"recovered" => 0.1,
"impossible" => 0.0,
"certain" => 1.0,
unexpected => panic!("unexpected state: {unexpected}"),
};
Json(json!({"answers": {"matches": {"type": "noul", "noul": probability}}}))
},
),
))
.await;
let ctx = context(server.function::<Noul>());
let events = "(VALUES (1, 'failed'), (2, 'boundary'), (3, 'recovered'), (4, NULL), (5, 'impossible'), (6, 'certain')) AS events(id, message)";
let batches = ctx
.sql(&format!(
"SELECT id FROM {events} WHERE ai_match((message), 'payment failed') >= 0.8 ORDER BY id"
))
.await
.unwrap()
.collect()
.await
.unwrap();
assert_batches_eq!(
[
"+----+", "| id |", "+----+", "| 1 |", "| 2 |", "| 6 |", "+----+"
],
&batches
);
let batches = ctx
.sql(&format!(
"SELECT id, ai_match(message, 'payment failed') AS score FROM {events} ORDER BY score DESC NULLS LAST"
))
.await
.unwrap()
.collect()
.await
.unwrap();
assert_eq!(batches[0].schema().field(1).data_type(), &DataType::Float64);
assert_batches_eq!(
[
"+----+-------+",
"| id | score |",
"+----+-------+",
"| 6 | 1.0 |",
"| 1 | 0.95 |",
"| 2 | 0.8 |",
"| 3 | 0.1 |",
"| 5 | 0.0 |",
"| 4 | |",
"+----+-------+"
],
&batches
);
}
#[tokio::test]
async fn test_ai_match_nulls_need_no_api_key() {
let ctx = context(AiFunction::<Noul> {
enabled: true,
api_key: None,
..Default::default()
});
let batches = ctx
.sql("SELECT ai_match(NULL, 'condition') AS a, ai_match('text', NULL) AS b")
.await
.unwrap()
.collect()
.await
.unwrap();
assert_eq!(batches[0].schema().field(0).data_type(), &DataType::Float64);
assert_batches_eq!(
[
"+---+---+",
"| a | b |",
"+---+---+",
"| | |",
"+---+---+"
],
&batches
);
let error = ctx
.sql("SELECT ai_match('text', 'condition')")
.await
.unwrap()
.collect()
.await
.unwrap_err();
assert!(error.to_string().contains("JEV_API_KEY"), "{error}");
let ctx = context(AiFunction::<Noul> {
enabled: false,
api_key: None,
..Default::default()
});
let error = ctx
.sql("SELECT ai_match('text', 'condition')")
.await
.unwrap()
.collect()
.await
.unwrap_err();
assert!(
error
.to_string()
.contains("GREPTIMEDB_EXPERIMENTAL_JEV=true"),
"{error}"
);
}
#[tokio::test]
async fn test_ai_match_bad_api_responses_fail_the_query() {
for (status, body, expected) in [
(
StatusCode::TOO_MANY_REQUESTS,
"{}",
"ai_match request failed",
),
(
StatusCode::OK,
"not json",
"ai_match response is not valid JSON",
),
(
StatusCode::OK,
r#"{"answers":{}}"#,
"ai_match response must contain",
),
(
StatusCode::OK,
r#"{"answers":{"matches":{"type":"noul","noul":1.1}}}"#,
"ai_match response must contain",
),
(
StatusCode::OK,
r#"{"answers":{"matches":{"type":"noul","noul":-0.1}}}"#,
"ai_match response must contain",
),
(
StatusCode::OK,
r#"{"answers":{"matches":{"type":"score","noul":0.9}}}"#,
"ai_match response must contain",
),
] {
let server = MockServer::start(
Router::new().route("/v1/systemone", post(move || async move { (status, body) })),
)
.await;
let ctx = context(server.function::<Noul>());
let error = ctx
.sql("SELECT ai_match('text', 'condition')")
.await
.unwrap()
.collect()
.await
.unwrap_err();
assert!(error.to_string().contains(expected), "{error}");
}
}
#[tokio::test]
async fn test_ai_choose_returns_and_filters_supplied_labels() {
let server = MockServer::start(Router::new().route(
"/v1/systemone",
post(|Json(request): Json<Value>| async move {
let question = &request["questions"]["matches"];
assert_eq!(question["type"], "choice");
assert_eq!(question["instructions"], "Route the ticket");
let choice = match request["state"].as_str().unwrap() {
"refund" => {
assert_eq!(question["criteria"], json!({"billing": null, "technical": {"examples": ["crash"]}}));
tokio::time::sleep(Duration::from_millis(20)).await;
"billing"
}
"crash" => {
assert_eq!(question["criteria"], json!({"billing": ["refunds"], "technical": "errors"}));
"technical"
}
unexpected => panic!("unexpected state: {unexpected}"),
};
Json(json!({"answers": {"matches": {"type": "choice", "choice": choice, "confidence": 0.7}}}))
}),
))
.await;
let ctx = context(server.function::<Choice>());
let events = r#"(VALUES
(1, 'refund', 'Route the ticket', '{"billing":null,"technical":{"examples":["crash"]}}'),
(2, 'crash', 'Route the ticket', '{"billing":["refunds"],"technical":"errors"}'),
(3, NULL, 'Route the ticket', 'invalid JSON'),
(4, 'refund', NULL, 'invalid JSON'),
(5, 'refund', 'Route the ticket', NULL)
) AS events(id, message, prompt, criteria)"#;
let batches = ctx
.sql(&format!(
"SELECT id, ai_choose(message, prompt, criteria) AS team FROM {events} ORDER BY id"
))
.await
.unwrap()
.collect()
.await
.unwrap();
assert_eq!(batches[0].schema().field(1).data_type(), &DataType::Utf8);
assert_batches_eq!(
[
"+----+-----------+",
"| id | team |",
"+----+-----------+",
"| 1 | billing |",
"| 2 | technical |",
"| 3 | |",
"| 4 | |",
"| 5 | |",
"+----+-----------+"
],
&batches
);
let batches = ctx
.sql(&format!(
"SELECT id FROM {events} WHERE ai_choose(message, prompt, criteria) = 'billing'"
))
.await
.unwrap()
.collect()
.await
.unwrap();
assert_batches_eq!(["+----+", "| id |", "+----+", "| 1 |", "+----+"], &batches);
let batches = ctx
.sql(
r#"SELECT id, ai_choose(message, 'Route the ticket',
'{"billing":null,"technical":{"examples":["crash"]}}') AS team
FROM (VALUES (1, 'refund'), (2, NULL), (3, 'refund')) AS events(id, message)
ORDER BY id"#,
)
.await
.unwrap()
.collect()
.await
.unwrap();
assert_batches_eq!(
[
"+----+---------+",
"| id | team |",
"+----+---------+",
"| 1 | billing |",
"| 2 | |",
"| 3 | billing |",
"+----+---------+"
],
&batches
);
}
#[tokio::test]
async fn test_ai_score_returns_fractional_scores_on_each_rows_scale() {
let server = MockServer::start(Router::new().route(
"/v1/systemone",
post(|Json(request): Json<Value>| async move {
let question = &request["questions"]["matches"];
assert_eq!(question["type"], "score");
assert_eq!(question["instructions"], "Rate severity");
let (score, probabilities) = match request["state"].as_str().unwrap() {
"low" => {
assert_eq!(question["criteria"], json!(["low", {"description": "medium"}, ["high"]]));
tokio::time::sleep(Duration::from_millis(20)).await;
(0.0, json!({"0":1.0,"1":0.0,"2":0.0}))
}
"fractional" => (1.25, json!({"2":0.25,"0":0.0,"1":0.75})),
"high" => (2.0, json!({"0":0.0,"1":0.0,"2":1.0})),
"binary" => {
assert_eq!(question["criteria"], json!(["low", "high"]));
(1.0, json!({"0":0.0,"1":1.0}))
}
"ten levels" => {
assert_eq!(question["criteria"].as_array().unwrap().len(), 10);
let probabilities: Map<String, Value> = (0..10)
.map(|level| (level.to_string(), json!(if level == 9 { 1.0 } else { 0.0 })))
.collect();
(9.0, Value::Object(probabilities))
}
unexpected => panic!("unexpected state: {unexpected}"),
};
Json(json!({"answers": {"matches": {"type": "score", "score": score, "confidence": 0.4, "probabilities": probabilities}}}))
}),
))
.await;
let ctx = context(server.function::<Score>());
let batches = ctx
.sql(
r#"
SELECT id, json_get_float(ai_score(message, 'Rate severity', criteria), 'score') AS severity
FROM (VALUES
(1, 'low', '["low", {"description":"medium"}, ["high"]]'),
(2, 'fractional', '["low", "medium", "high"]'),
(3, 'high', '["low", "medium", "high"]'),
(4, 'binary', '["low", "high"]'),
(5, 'ten levels', '["a","b","c","d","e","f","g","h","i","j"]'),
(6, NULL, 'invalid JSON'),
(7, 'high', NULL)
) AS events(id, message, criteria)
ORDER BY severity DESC NULLS LAST, id
"#,
)
.await
.unwrap()
.collect()
.await
.unwrap();
assert_eq!(batches[0].schema().field(1).data_type(), &DataType::Float64);
assert_batches_eq!(
[
"+----+----------+",
"| id | severity |",
"+----+----------+",
"| 5 | 9.0 |",
"| 3 | 2.0 |",
"| 2 | 1.25 |",
"| 4 | 1.0 |",
"| 1 | 0.0 |",
"| 6 | |",
"| 7 | |",
"+----+----------+"
],
&batches
);
let batches = ctx
.sql(
r#"SELECT id, json_get_float(ai_score(message, 'Rate severity',
'["low",{"description":"medium"},["high"]]'), 'score') AS severity
FROM (VALUES (1, 'low'), (2, 'fractional'), (3, 'high'), (4, NULL)) AS events(id, message)
ORDER BY id"#,
)
.await
.unwrap()
.collect()
.await
.unwrap();
assert_batches_eq!(
[
"+----+----------+",
"| id | severity |",
"+----+----------+",
"| 1 | 0.0 |",
"| 2 | 1.25 |",
"| 3 | 2.0 |",
"| 4 | |",
"+----+----------+"
],
&batches
);
}
#[tokio::test]
async fn test_ai_score_preserves_uncertainty_in_one_evaluation() {
let calls = Arc::new(AtomicUsize::new(0));
let request_calls = calls.clone();
let server = MockServer::start(Router::new().route(
"/v1/systemone",
post(move |Json(request): Json<Value>| {
request_calls.fetch_add(1, Ordering::Relaxed);
async move {
let (confidence, probabilities) = match request["state"].as_str().unwrap() {
"medium" => (1.0, json!({"2":0.0,"0":0.0,"1":1.0})),
"split" => (0.25, json!({"2":0.5,"0":0.5,"1":0.0})),
unexpected => panic!("unexpected state: {unexpected}"),
};
Json(json!({"answers":{"matches":{
"type":"score", "score":1.0, "confidence":confidence,
"probabilities":probabilities
}}}))
}
}),
))
.await;
let ctx = context(server.function::<Score>());
let rated = r#"(SELECT id, ai_score(message, 'Rate severity', '["low","medium","high"]') AS rating
FROM (VALUES (1, 'medium'), (2, 'split'), (3, NULL)) AS events(id, message)) AS rated"#;
let batches = ctx
.sql(&format!(
"SELECT id, json_get_float(rating, 'score') AS score,
json_get_float(rating, 'confidence') AS confidence,
json_get_float(rating, 'probabilities[0]') AS p0,
json_get_float(rating, 'probabilities[1]') AS p1,
json_get_float(rating, 'probabilities[2]') AS p2
FROM {rated} ORDER BY id"
))
.await
.unwrap()
.collect()
.await
.unwrap();
assert_batches_eq!(
[
"+----+-------+------------+-----+-----+-----+",
"| id | score | confidence | p0 | p1 | p2 |",
"+----+-------+------------+-----+-----+-----+",
"| 1 | 1.0 | 1.0 | 0.0 | 1.0 | 0.0 |",
"| 2 | 1.0 | 0.25 | 0.5 | 0.0 | 0.5 |",
"| 3 | | | | | |",
"+----+-------+------------+-----+-----+-----+"
],
&batches
);
assert_eq!(calls.load(Ordering::Relaxed), 2);
let batches = ctx
.sql(&format!(
"SELECT id FROM {rated} WHERE json_get_float(rating, 'confidence') >= 0.8
ORDER BY json_get_float(rating, 'score') DESC"
))
.await
.unwrap()
.collect()
.await
.unwrap();
assert_batches_eq!(["+----+", "| id |", "+----+", "| 1 |", "+----+"], &batches);
assert_eq!(calls.load(Ordering::Relaxed), 4);
}
#[test]
fn test_ai_score_preserves_provider_precision() {
let answer = json!({
"score":1.6, "confidence":0.92,
"probabilities":{"2":0.6999995,"0":0.1,"1":0.2}
});
let criteria = vec![json!("low"), json!("medium"), json!("high")];
let ScalarValue::BinaryView(Some(bytes)) = Score::parse_answer(&answer, &criteria).unwrap()
else {
panic!("expected a JSONB object");
};
let object: Value =
serde_json::from_str(&jsonb::from_slice(&bytes).unwrap().to_string()).unwrap();
assert_eq!(object.as_object().unwrap().len(), 3);
assert_eq!(object["score"].as_f64(), Some(1.6));
assert_eq!(object["confidence"].as_f64(), Some(0.92));
let probabilities: Vec<f64> = object["probabilities"]
.as_array()
.unwrap()
.iter()
.map(|probability| probability.as_f64().unwrap())
.collect();
assert_eq!(probabilities, vec![0.1, 0.2, 0.6999995]);
}
#[tokio::test]
async fn test_ai_score_rejects_invalid_uncertainty() {
for (field, value) in [
("confidence", Value::Null),
("confidence", json!(-0.1)),
("confidence", json!(1.1)),
("confidence", json!("0.9")),
("probabilities", Value::Null),
("probabilities", json!([0.0, 1.0, 0.0])),
("probabilities", json!({"0":1.0})),
("probabilities", json!({"0":0.5,"2":0.5,"3":0.0})),
("probabilities", json!({"0":0.0,"1":1.0,"2":0.0,"3":0.0})),
("probabilities", json!({"0":-0.1,"1":0.6,"2":0.5})),
("probabilities", json!({"0":0.0,"1":1.1,"2":0.0})),
("probabilities", json!({"0":0.0,"1":0.95,"2":"0.05"})),
("probabilities", json!({"0":0.0,"1":0.0,"2":0.0})),
("probabilities", json!({"0":0.5,"1":0.5,"2":0.5})),
] {
let mut answer = json!({
"type":"score", "score":1.0, "confidence":0.9,
"probabilities":{"0":0.0,"1":1.0,"2":0.0}
});
answer[field] = value;
let server = MockServer::start(Router::new().route(
"/v1/systemone",
post(move || {
let answer = answer.clone();
async move { Json(json!({"answers":{"matches":answer}})) }
}),
))
.await;
let ctx = context(server.function::<Score>());
let error = ctx
.sql(r#"SELECT ai_score('text', 'prompt', '["low","medium","high"]')"#)
.await
.unwrap()
.collect()
.await
.unwrap_err();
assert!(error.to_string().contains("ai_score response"), "{error}");
}
}
#[tokio::test]
async fn test_ai_invalid_criteria_fail_before_requests() {
let choice_ctx = context(AiFunction::<Choice> {
enabled: true,
api_key: None,
..Default::default()
});
let score_ctx = context(AiFunction::<Score> {
enabled: true,
api_key: None,
..Default::default()
});
let too_many_choices: Map<String, Value> = (0..256)
.map(|index| (index.to_string(), Value::Null))
.collect();
for (ctx, name, valid, invalid) in [
(
&choice_ctx,
"ai_choose",
r#"{"billing":null}"#,
vec![
"not json".to_string(),
"null".to_string(),
"{}".to_string(),
"[]".to_string(),
r#"{"billing":1}"#.to_string(),
r#"{"billing":true}"#.to_string(),
Value::Object(too_many_choices).to_string(),
],
),
(
&score_ctx,
"ai_score",
r#"["low","high"]"#,
vec![
"not json".to_string(),
"null".to_string(),
"{}".to_string(),
"[]".to_string(),
r#"["low"]"#.to_string(),
r#"["low",null]"#.to_string(),
r#"["low",1]"#.to_string(),
r#"["low",true]"#.to_string(),
json!(vec!["level"; 11]).to_string(),
],
),
] {
let error = ctx
.sql(&format!(
"SELECT {name}(message, 'prompt', 'not json') FROM (VALUES (NULL), ('text')) AS input(message)"
))
.await
.unwrap()
.collect()
.await
.unwrap_err();
assert!(
error.to_string().contains(&format!("{name} criteria")),
"{error}"
);
for criteria in invalid {
// A later invalid row must fail before even creating a client for the first row.
let error = ctx.sql(&format!(
"SELECT {name}('text', 'prompt', criteria) FROM (VALUES ('{valid}'), ('{criteria}')) AS input(criteria)"
)).await.unwrap().collect().await.unwrap_err();
assert!(
error.to_string().contains(&format!("{name} criteria")),
"{error}"
);
}
}
}
#[tokio::test]
async fn test_ai_choose_and_score_reject_invalid_answers() {
for (name, criteria, answer) in [
(
"ai_choose",
r#"{"billing":null}"#,
json!({"type":"score","choice":"billing"}),
),
(
"ai_choose",
r#"{"billing":null}"#,
json!({"type":"choice","choice":"unknown"}),
),
(
"ai_choose",
r#"{"billing":null}"#,
json!({"type":"choice","choice":1}),
),
("ai_choose", r#"{"billing":null}"#, json!({"type":"choice"})),
(
"ai_score",
r#"["low","medium","high"]"#,
json!({"type":"noul","score":0.5}),
),
(
"ai_score",
r#"["low","medium","high"]"#,
json!({"type":"score","score":2.1}),
),
(
"ai_score",
r#"["low","high"]"#,
json!({"type":"score","score":-0.1}),
),
(
"ai_score",
r#"["low","high"]"#,
json!({"type":"score","score":"NaN"}),
),
("ai_score", r#"["low","high"]"#, json!({"type":"score"})),
] {
let server = MockServer::start(Router::new().route(
"/v1/systemone",
post(move || {
let answer = answer.clone();
async move { Json(json!({"answers":{"matches":answer}})) }
}),
))
.await;
let ctx = context(server.function::<Choice>());
ctx.register_udf(
AsyncScalarUDF::new(Arc::new(server.function::<Score>())).into_scalar_udf(),
);
let error = ctx
.sql(&format!("SELECT {name}('text', 'prompt', '{criteria}')"))
.await
.unwrap()
.collect()
.await
.unwrap_err();
assert!(
error
.to_string()
.contains(&format!("{name} response must contain")),
"{error}"
);
}
}
#[tokio::test]
#[ignore = "calls the real Jev API; requires JEV_API_KEY and GREPTIMEDB_EXPERIMENTAL_JEV=true"]
async fn test_ai_match_live() {
let ctx = SessionContext::new();
let function = FUNCTION_REGISTRY.get_function("ai_match").unwrap();
ctx.register_udf(function.provide(FunctionContext::default()));
let batches = ctx.sql(
"SELECT id FROM (VALUES
(1, 'Payment failed permanently: all three retries exhausted; the payment is still unsuccessful.'),
(2, 'Payment succeeded on the second retry; the payment is complete.'),
(3, 'User logged in successfully.')
) AS events(id, message)
WHERE ai_match(message, 'The event reports that a payment still failed after retries.') >= 0.8
ORDER BY id",
).await.unwrap().collect().await.unwrap();
assert_batches_eq!(["+----+", "| id |", "+----+", "| 1 |", "+----+"], &batches);
}
-209
View File
@@ -1,209 +0,0 @@
// Copyright 2023 Greptime Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#[cfg(test)]
mod tests;
use std::fmt;
use std::sync::Arc;
use std::time::Duration;
use arrow::array::BooleanArray;
use arrow::datatypes::DataType;
use async_trait::async_trait;
use datafusion_common::cast::{as_float64_array, as_string_view_array};
use datafusion_common::utils::take_function_args;
use datafusion_common::{Result, exec_datafusion_err, exec_err, not_impl_err};
use datafusion_expr::async_udf::{AsyncScalarUDF, AsyncScalarUDFImpl};
use datafusion_expr::{ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl, Signature, Volatility};
use futures::{StreamExt, TryStreamExt, stream};
use reqwest::Client;
use serde_json::{Value, json};
use crate::function_factory::ScalarFunctionFactory;
use crate::function_registry::FunctionRegistry;
/// Experimental natural-language predicate: `jev(text, statement, threshold)`.
/// Each non-null row is a Noul question; its probability is compared with `>=`.
#[derive(PartialEq, Eq, Hash)]
pub(crate) struct JevFunction {
signature: Signature,
enabled: bool,
api_key: Option<String>,
endpoint: String,
model: String,
}
impl JevFunction {
pub(crate) fn register(registry: &FunctionRegistry) {
registry.register(ScalarFunctionFactory {
name: "jev".to_string(),
factory: Arc::new(|_| AsyncScalarUDF::new(Arc::new(Self::default())).into_scalar_udf()),
});
}
async fn evaluate(&self, client: &Client, text: &str, statement: &str) -> Result<f64> {
let response: Value = client
.post(&self.endpoint)
.json(&json!({
"model": self.model,
"state": text,
"questions": {
"matches": { "type": "noul", "instructions": statement }
}
}))
.send()
.await
.and_then(|response| response.error_for_status())
.map_err(|e| exec_datafusion_err!("jev request failed: {e}"))?
.json()
.await
.map_err(|e| exec_datafusion_err!("jev response is not valid JSON: {e}"))?;
let answer = &response["answers"]["matches"];
let probability = answer["noul"].as_f64().filter(|p| (0.0..=1.0).contains(p));
match (answer["type"].as_str(), probability) {
(Some("noul"), Some(probability)) => Ok(probability),
_ => exec_err!(
"jev response must contain answers.matches with type noul and a probability in [0, 1]"
),
}
}
fn client(&self) -> Result<Client> {
if !self.enabled {
return exec_err!(
"jev is experimental; set GREPTIMEDB_EXPERIMENTAL_JEV=true to enable it"
);
}
let key = self
.api_key
.as_deref()
.filter(|key| !key.trim().is_empty())
.ok_or_else(|| {
exec_datafusion_err!("jev requires the JEV_API_KEY environment variable")
})?;
let mut authorization = reqwest::header::HeaderValue::from_str(&format!("Bearer {key}"))
.map_err(|_| exec_datafusion_err!("JEV_API_KEY is not a valid HTTP header value"))?;
authorization.set_sensitive(true);
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(reqwest::header::AUTHORIZATION, authorization);
Client::builder()
.default_headers(headers)
.timeout(Duration::from_secs(30))
.build()
.map_err(|e| exec_datafusion_err!("failed to create jev HTTP client: {e}"))
}
}
impl Default for JevFunction {
fn default() -> Self {
Self {
// External model evaluations must not be constant-folded during planning.
signature: Signature::exact(
vec![DataType::Utf8View, DataType::Utf8View, DataType::Float64],
Volatility::Volatile,
),
enabled: std::env::var("GREPTIMEDB_EXPERIMENTAL_JEV").as_deref() == Ok("true"),
api_key: std::env::var("JEV_API_KEY").ok(),
endpoint: std::env::var("JEV_ENDPOINT")
.unwrap_or_else(|_| "https://api.typesafe.ai/v1/systemone".to_string()),
model: std::env::var("JEV_MODEL").unwrap_or_else(|_| "jev-latest".to_string()),
}
}
}
impl fmt::Debug for JevFunction {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
// Credentials must never appear in plans or diagnostic output.
f.debug_struct("JevFunction")
.field("signature", &self.signature)
.field("enabled", &self.enabled)
.field("model", &self.model)
.finish_non_exhaustive()
}
}
impl ScalarUDFImpl for JevFunction {
fn name(&self) -> &str {
"jev"
}
fn signature(&self) -> &Signature {
&self.signature
}
fn return_type(&self, _arg_types: &[DataType]) -> Result<DataType> {
Ok(DataType::Boolean)
}
fn invoke_with_args(&self, _args: ScalarFunctionArgs) -> Result<ColumnarValue> {
not_impl_err!("jev can only be called from async contexts")
}
}
#[async_trait]
impl AsyncScalarUDFImpl for JevFunction {
async fn invoke_async_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
let arrays = args
.args
.into_iter()
.map(|arg| arg.into_array(args.number_rows))
.collect::<Result<Vec<_>>>()?;
let [texts, statements, thresholds] = take_function_args(self.name(), arrays)?;
let texts = as_string_view_array(&texts)?;
let statements = as_string_view_array(&statements)?;
let thresholds = as_float64_array(&thresholds)?;
let rows: Vec<_> = texts
.iter()
.zip(statements.iter())
.zip(thresholds.iter())
.map(|((text, statement), threshold)| Some((text?, statement?, threshold?)))
.collect();
// Validate the whole batch before making any billable requests.
for (_, _, threshold) in rows.iter().flatten() {
if !(0.0..=1.0).contains(threshold) {
return exec_err!("jev threshold must be finite and in [0, 1]");
}
}
if rows.iter().all(Option::is_none) {
return Ok(ColumnarValue::Array(Arc::new(BooleanArray::new_null(
args.number_rows,
))));
}
// Reuse connections within the batch; bound in-flight requests and preserve row order.
let client = self.client()?;
let requests: Vec<_> = rows
.into_iter()
.map(|row| {
let client = &client;
async move {
match row {
Some((text, statement, threshold)) => self
.evaluate(client, text, statement)
.await
.map(|p| Some(p >= threshold)),
None => Ok(None),
}
}
})
.collect();
// This limit is per expression/batch invocation, not per query or process.
// Concurrent partitions and queries can each have their own in-flight requests.
let matches: Vec<Option<bool>> = stream::iter(requests).buffered(8).try_collect().await?;
Ok(ColumnarValue::Array(Arc::new(BooleanArray::from(matches))))
}
}
@@ -1,250 +0,0 @@
// Copyright 2023 Greptime Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use axum::http::{HeaderMap, StatusCode};
use axum::routing::post;
use axum::{Json, Router};
use datafusion::assert_batches_eq;
use datafusion::prelude::{SessionConfig, SessionContext};
use tokio::net::TcpListener;
use tokio::task::JoinHandle;
use super::*;
use crate::function::FunctionContext;
use crate::function_registry::FUNCTION_REGISTRY;
struct MockServer {
endpoint: String,
task: JoinHandle<()>,
}
impl MockServer {
async fn start(router: Router) -> Self {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let endpoint = format!("http://{}/v1/systemone", listener.local_addr().unwrap());
let task = tokio::spawn(async move { axum::serve(listener, router).await.unwrap() });
Self { endpoint, task }
}
fn function(&self) -> JevFunction {
JevFunction {
enabled: true,
api_key: Some("test-key".to_string()),
endpoint: self.endpoint.clone(),
model: "test-model".to_string(),
..Default::default()
}
}
}
impl Drop for MockServer {
fn drop(&mut self) {
self.task.abort();
}
}
fn context(function: JevFunction) -> SessionContext {
let ctx = SessionContext::new_with_config(SessionConfig::new().with_target_partitions(1));
ctx.register_udf(AsyncScalarUDF::new(Arc::new(function)).into_scalar_udf());
ctx
}
#[tokio::test]
async fn test_jev_sql_filters_using_noul_probability() {
let server = MockServer::start(Router::new().route(
"/v1/systemone",
post(
|headers: HeaderMap, Json(request): Json<Value>| async move {
assert_eq!(headers["authorization"], "Bearer test-key");
assert_eq!(request["model"], "test-model");
assert_eq!(request["questions"]["matches"]["type"], "noul");
// The SQL statement is passed through verbatim, not rewritten as a prompt.
assert_eq!(
request["questions"]["matches"]["instructions"],
"payment failed"
);
let probability = match request["state"].as_str().unwrap() {
"failed" => {
// Complete this row after later rows to exercise result alignment.
tokio::time::sleep(Duration::from_millis(20)).await;
0.95
}
"boundary" => 0.8,
"recovered" => 0.1,
unexpected => panic!("unexpected state: {unexpected}"),
};
Json(json!({"answers": {"matches": {"type": "noul", "noul": probability}}}))
},
),
))
.await;
let ctx = context(server.function());
let events = "(VALUES (1, 'failed'), (2, 'boundary'), (3, 'recovered'), (4, NULL)) AS events(id, message)";
let batches = ctx
.sql(&format!(
"SELECT id FROM {events} WHERE jev((message), 'payment failed', 0.8) ORDER BY id"
))
.await
.unwrap()
.collect()
.await
.unwrap();
assert_batches_eq!(
["+----+", "| id |", "+----+", "| 1 |", "| 2 |", "+----+"],
&batches
);
let batches = ctx
.sql(&format!(
"SELECT id, jev(message, 'payment failed', 0.8) AS matched FROM {events} ORDER BY id"
))
.await
.unwrap()
.collect()
.await
.unwrap();
assert_batches_eq!(
[
"+----+---------+",
"| id | matched |",
"+----+---------+",
"| 1 | true |",
"| 2 | true |",
"| 3 | false |",
"| 4 | |",
"+----+---------+"
],
&batches
);
}
#[tokio::test]
async fn test_jev_nulls_and_invalid_thresholds_need_no_api_key() {
let ctx = context(JevFunction {
enabled: true,
api_key: None,
..Default::default()
});
let batches = ctx
.sql("SELECT jev(NULL, 'condition', 0.8) AS a, jev('text', NULL, 0.8) AS b, jev('text', 'condition', NULL) AS c")
.await.unwrap().collect().await.unwrap();
assert_batches_eq!(
[
"+---+---+---+",
"| a | b | c |",
"+---+---+---+",
"| | | |",
"+---+---+---+"
],
&batches
);
for threshold in ["-0.1", "1.1", "'NaN'::DOUBLE", "'Infinity'::DOUBLE"] {
let error = ctx
.sql(&format!("SELECT jev('text', 'condition', {threshold})"))
.await
.unwrap()
.collect()
.await
.unwrap_err();
assert!(
error
.to_string()
.contains("jev threshold must be finite and in [0, 1]"),
"{error}"
);
}
let error = ctx
.sql("SELECT jev('text', 'condition', 0.8)")
.await
.unwrap()
.collect()
.await
.unwrap_err();
assert!(error.to_string().contains("JEV_API_KEY"), "{error}");
let ctx = context(JevFunction {
enabled: false,
api_key: None,
..Default::default()
});
let error = ctx
.sql("SELECT jev('text', 'condition', 0.8)")
.await
.unwrap()
.collect()
.await
.unwrap_err();
assert!(
error
.to_string()
.contains("GREPTIMEDB_EXPERIMENTAL_JEV=true"),
"{error}"
);
}
#[tokio::test]
async fn test_jev_bad_api_responses_fail_the_query() {
for (status, body, expected) in [
(StatusCode::TOO_MANY_REQUESTS, "{}", "jev request failed"),
(StatusCode::OK, "not json", "jev response is not valid JSON"),
(
StatusCode::OK,
r#"{"answers":{}}"#,
"jev response must contain",
),
(
StatusCode::OK,
r#"{"answers":{"matches":{"type":"noul","noul":1.1}}}"#,
"jev response must contain",
),
(
StatusCode::OK,
r#"{"answers":{"matches":{"type":"score","noul":0.9}}}"#,
"jev response must contain",
),
] {
let server = MockServer::start(
Router::new().route("/v1/systemone", post(move || async move { (status, body) })),
)
.await;
let ctx = context(server.function());
let error = ctx
.sql("SELECT jev('text', 'condition', 0.8)")
.await
.unwrap()
.collect()
.await
.unwrap_err();
assert!(error.to_string().contains(expected), "{error}");
}
}
#[tokio::test]
#[ignore = "calls the real Jev API; requires JEV_API_KEY and GREPTIMEDB_EXPERIMENTAL_JEV=true"]
async fn test_jev_live() {
let ctx = SessionContext::new();
let function = FUNCTION_REGISTRY.get_function("jev").unwrap();
ctx.register_udf(function.provide(FunctionContext::default()));
let batches = ctx.sql(
"SELECT id FROM (VALUES
(1, 'Payment failed permanently: all three retries exhausted; the payment is still unsuccessful.'),
(2, 'Payment succeeded on the second retry; the payment is complete.'),
(3, 'User logged in successfully.')
) AS events(id, message)
WHERE jev(message, 'The event reports that a payment still failed after retries.', 0.8)
ORDER BY id",
).await.unwrap().collect().await.unwrap();
assert_batches_eq!(["+----+", "| id |", "+----+", "| 1 |", "+----+"], &batches);
}
@@ -0,0 +1,144 @@
-- No credentials or external service needed: NULL propagates without an API call.
SELECT ai_match(NULL, 'The event reports a failed payment.') AS score;
+-------+
| score |
+-------+
| |
+-------+
SELECT ai_match('message', NULL) AS score;
+-------+
| score |
+-------+
| |
+-------+
SELECT arrow_typeof(ai_match(NULL, 'condition')) AS score_type;
+------------+
| score_type |
+------------+
| Float64 |
+------------+
SELECT ai_choose(NULL, 'Route the ticket', '{"billing":null}') AS null_text,
ai_choose('message', NULL, 'invalid JSON') AS null_prompt,
ai_choose('message', 'Route the ticket', NULL) AS null_criteria;
+-----------+-------------+---------------+
| null_text | null_prompt | null_criteria |
+-----------+-------------+---------------+
| | | |
+-----------+-------------+---------------+
SELECT ai_score(NULL, 'Rate severity', '["low","high"]') AS null_text,
ai_score('message', NULL, 'invalid JSON') AS null_prompt,
ai_score('message', 'Rate severity', NULL) AS null_criteria;
+-----------+-------------+---------------+
| null_text | null_prompt | null_criteria |
+-----------+-------------+---------------+
| | | |
+-----------+-------------+---------------+
SELECT arrow_typeof(ai_choose(NULL, 'prompt', '{"billing":null}')) AS choice_type,
arrow_typeof(ai_score(NULL, 'prompt', '["low","high"]')) AS score_type;
+-------------+------------+
| choice_type | score_type |
+-------------+------------+
| Utf8 | BinaryView |
+-------------+------------+
SELECT json_to_string(rating) AS rating,
json_get_float(rating, 'score') AS score,
json_get_float(rating, 'confidence') AS confidence,
json_get_float(rating, 'probabilities[2]') AS high_probability
FROM (SELECT ai_score(NULL, 'prompt', '["low","medium","high"]') AS rating) AS rated;
+--------+-------+------------+------------------+
| rating | score | confidence | high_probability |
+--------+-------+------------+------------------+
| | | | |
+--------+-------+------------+------------------+
-- Validate SQL registration, coercion, and asynchronous filtering on a real table.
CREATE TABLE ai_events (
occurred_at TIMESTAMP TIME INDEX,
"service" STRING,
"message" STRING,
PRIMARY KEY ("service")
);
Affected Rows: 0
INSERT INTO ai_events VALUES
('2026-09-19T01:00:00Z', 'payments', NULL),
('2026-09-19T02:00:00Z', 'auth', NULL);
Affected Rows: 2
SELECT occurred_at, service, message FROM ai_events
WHERE occurred_at >= '2026-09-19T00:00:00Z'
AND occurred_at < '2026-09-20T00:00:00Z'
AND service = 'payments'
AND ai_match((message), 'The event reports that a payment still failed after retries.') >= 0.8
ORDER BY occurred_at;
+-------------+---------+---------+
| occurred_at | service | message |
+-------------+---------+---------+
+-------------+---------+---------+
SELECT occurred_at,
ai_choose(message, 'Route the ticket', '{"billing":"Payments","technical":"Errors"}') AS team,
ai_score(message, 'Rate severity', '["low","medium","high"]') AS rating
FROM ai_events
ORDER BY occurred_at;
+---------------------+------+--------+
| occurred_at | team | rating |
+---------------------+------+--------+
| 2026-09-19T01:00:00 | | |
| 2026-09-19T02:00:00 | | |
+---------------------+------+--------+
DROP TABLE ai_events;
Affected Rows: 0
-- Matching uses two arguments; thresholds are expressed as SQL comparisons.
SELECT ai_match('message', 'condition', 0.8);
Error: 3000(PlanQuery), Failed to plan SQL: Error during planning: Failed to coerce arguments to satisfy a call to 'ai_match' function: coercion from Utf8, Utf8, Float64 to the signature Exact(Utf8View, Utf8View) failed. No function matches the given name and argument types 'ai_match(Utf8, Utf8, Float64)'. You might need to add explicit type casts.
Candidate functions:
ai_match(Utf8View, Utf8View)
-- Invalid criteria fail locally before any HTTP request.
SELECT ai_choose('message', 'prompt', 'not json');
Error: 3001(EngineExecuteQuery), Execution error: ai_choose criteria is not valid JSON: expected ident at line 1 column 2
SELECT ai_choose('message', 'prompt', '{}');
Error: 3001(EngineExecuteQuery), Execution error: ai_choose criteria must be a JSON object with 1 to 255 options; descriptions must be strings, objects, arrays, or null
SELECT ai_score('message', 'prompt', '["only one level"]');
Error: 3001(EngineExecuteQuery), Execution error: ai_score criteria must be a JSON array with 2 to 10 levels; descriptions must be strings, objects, or arrays
-- Both additional modes require criteria.
SELECT ai_choose('message', 'prompt');
Error: 3000(PlanQuery), Failed to plan SQL: Error during planning: Failed to coerce arguments to satisfy a call to 'ai_choose' function: coercion from Utf8, Utf8 to the signature Exact(Utf8View, Utf8View, Utf8View) failed. No function matches the given name and argument types 'ai_choose(Utf8, Utf8)'. You might need to add explicit type casts.
Candidate functions:
ai_choose(Utf8View, Utf8View, Utf8View)
SELECT ai_score('message', 'prompt');
Error: 3000(PlanQuery), Failed to plan SQL: Error during planning: Failed to coerce arguments to satisfy a call to 'ai_score' function: coercion from Utf8, Utf8 to the signature Exact(Utf8View, Utf8View, Utf8View) failed. No function matches the given name and argument types 'ai_score(Utf8, Utf8)'. You might need to add explicit type casts.
Candidate functions:
ai_score(Utf8View, Utf8View, Utf8View)
@@ -0,0 +1,65 @@
-- No credentials or external service needed: NULL propagates without an API call.
SELECT ai_match(NULL, 'The event reports a failed payment.') AS score;
SELECT ai_match('message', NULL) AS score;
SELECT arrow_typeof(ai_match(NULL, 'condition')) AS score_type;
SELECT ai_choose(NULL, 'Route the ticket', '{"billing":null}') AS null_text,
ai_choose('message', NULL, 'invalid JSON') AS null_prompt,
ai_choose('message', 'Route the ticket', NULL) AS null_criteria;
SELECT ai_score(NULL, 'Rate severity', '["low","high"]') AS null_text,
ai_score('message', NULL, 'invalid JSON') AS null_prompt,
ai_score('message', 'Rate severity', NULL) AS null_criteria;
SELECT arrow_typeof(ai_choose(NULL, 'prompt', '{"billing":null}')) AS choice_type,
arrow_typeof(ai_score(NULL, 'prompt', '["low","high"]')) AS score_type;
SELECT json_to_string(rating) AS rating,
json_get_float(rating, 'score') AS score,
json_get_float(rating, 'confidence') AS confidence,
json_get_float(rating, 'probabilities[2]') AS high_probability
FROM (SELECT ai_score(NULL, 'prompt', '["low","medium","high"]') AS rating) AS rated;
-- Validate SQL registration, coercion, and asynchronous filtering on a real table.
CREATE TABLE ai_events (
occurred_at TIMESTAMP TIME INDEX,
"service" STRING,
"message" STRING,
PRIMARY KEY ("service")
);
INSERT INTO ai_events VALUES
('2026-09-19T01:00:00Z', 'payments', NULL),
('2026-09-19T02:00:00Z', 'auth', NULL);
SELECT occurred_at, service, message FROM ai_events
WHERE occurred_at >= '2026-09-19T00:00:00Z'
AND occurred_at < '2026-09-20T00:00:00Z'
AND service = 'payments'
AND ai_match((message), 'The event reports that a payment still failed after retries.') >= 0.8
ORDER BY occurred_at;
SELECT occurred_at,
ai_choose(message, 'Route the ticket', '{"billing":"Payments","technical":"Errors"}') AS team,
ai_score(message, 'Rate severity', '["low","medium","high"]') AS rating
FROM ai_events
ORDER BY occurred_at;
DROP TABLE ai_events;
-- Matching uses two arguments; thresholds are expressed as SQL comparisons.
SELECT ai_match('message', 'condition', 0.8);
-- Invalid criteria fail locally before any HTTP request.
SELECT ai_choose('message', 'prompt', 'not json');
SELECT ai_choose('message', 'prompt', '{}');
SELECT ai_score('message', 'prompt', '["only one level"]');
-- Both additional modes require criteria.
SELECT ai_choose('message', 'prompt');
SELECT ai_score('message', 'prompt');
@@ -1,66 +0,0 @@
-- No credentials or external service needed: NULL propagates without an API call.
SELECT jev(NULL, 'The event reports a failed payment.', 0.8) AS matched;
+---------+
| matched |
+---------+
| |
+---------+
SELECT jev('message', NULL, 0.8) AS matched;
+---------+
| matched |
+---------+
| |
+---------+
SELECT jev('message', 'The event reports a failed payment.', NULL) AS matched;
+---------+
| matched |
+---------+
| |
+---------+
-- Validate SQL registration, coercion, and asynchronous filtering on a real table.
CREATE TABLE jev_events (
occurred_at TIMESTAMP TIME INDEX,
"service" STRING,
"message" STRING,
PRIMARY KEY ("service")
);
Affected Rows: 0
INSERT INTO jev_events VALUES
('2026-09-19T01:00:00Z', 'payments', NULL),
('2026-09-19T02:00:00Z', 'auth', NULL);
Affected Rows: 2
SELECT occurred_at, service, message FROM jev_events
WHERE occurred_at >= '2026-09-19T00:00:00Z'
AND occurred_at < '2026-09-20T00:00:00Z'
AND service = 'payments'
AND jev((message), 'The event reports that a payment still failed after retries.', 0.8)
ORDER BY occurred_at;
+-------------+---------+---------+
| occurred_at | service | message |
+-------------+---------+---------+
+-------------+---------+---------+
DROP TABLE jev_events;
Affected Rows: 0
-- Invalid thresholds fail locally before any HTTP request.
SELECT jev('message', 'condition', -0.1);
Error: 3001(EngineExecuteQuery), Execution error: jev threshold must be finite and in [0, 1]
SELECT jev('message', 'condition', 1.1);
Error: 3001(EngineExecuteQuery), Execution error: jev threshold must be finite and in [0, 1]
@@ -1,32 +0,0 @@
-- No credentials or external service needed: NULL propagates without an API call.
SELECT jev(NULL, 'The event reports a failed payment.', 0.8) AS matched;
SELECT jev('message', NULL, 0.8) AS matched;
SELECT jev('message', 'The event reports a failed payment.', NULL) AS matched;
-- Validate SQL registration, coercion, and asynchronous filtering on a real table.
CREATE TABLE jev_events (
occurred_at TIMESTAMP TIME INDEX,
"service" STRING,
"message" STRING,
PRIMARY KEY ("service")
);
INSERT INTO jev_events VALUES
('2026-09-19T01:00:00Z', 'payments', NULL),
('2026-09-19T02:00:00Z', 'auth', NULL);
SELECT occurred_at, service, message FROM jev_events
WHERE occurred_at >= '2026-09-19T00:00:00Z'
AND occurred_at < '2026-09-20T00:00:00Z'
AND service = 'payments'
AND jev((message), 'The event reports that a payment still failed after retries.', 0.8)
ORDER BY occurred_at;
DROP TABLE jev_events;
-- Invalid thresholds fail locally before any HTTP request.
SELECT jev('message', 'condition', -0.1);
SELECT jev('message', 'condition', 1.1);