diff --git a/docs/how-to/ai-functions.md b/docs/how-to/ai-functions.md new file mode 100644 index 00000000000..5c33ffff50e --- /dev/null +++ b/docs/how-to/ai-functions.md @@ -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='' +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: + +## 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 +``` diff --git a/docs/how-to/jev.md b/docs/how-to/jev.md deleted file mode 100644 index 2a21f085946..00000000000 --- a/docs/how-to/jev.md +++ /dev/null @@ -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='' -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: - -## 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 -``` diff --git a/src/common/function/src/function_registry.rs b/src/common/function/src/function_registry.rs index 037ca5b175a..6583f7ba607 100644 --- a/src/common/function/src/function_registry.rs +++ b/src/common/function/src/function_registry.rs @@ -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> = 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] diff --git a/src/common/function/src/scalars.rs b/src/common/function/src/scalars.rs index 61fb7ec632c..4bc4c17f5bb 100644 --- a/src/common/function/src/scalars.rs +++ b/src/common/function/src/scalars.rs @@ -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; diff --git a/src/common/function/src/scalars/ai.rs b/src/common/function/src/scalars/ai.rs new file mode 100644 index 00000000000..f2e17dffd36 --- /dev/null +++ b/src/common/function/src/scalars/ai.rs @@ -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::::register(registry); + AiFunction::::register(registry); + AiFunction::::register(registry); +} + +/// Shared asynchronous execution for AI functions using the Jev backend. +#[derive(PartialEq, Eq, Hash)] +struct AiFunction { + signature: Signature, + question: PhantomData, + enabled: bool, + api_key: Option, + endpoint: String, + model: String, +} + +struct AiRequest<'a, Q: AiQuestion> { + text: &'a str, + prompt: &'a str, + criteria: Arc, +} + +/// 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 { + 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 AiFunction { + 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 { + 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>>> { + 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::>>()?; + let scalar_criteria = args[2..] + .iter() + .all(|arg| matches!(arg, StringArgument::Scalar(_))); + let mut shared_criteria: Option> = None; + + // Validate all non-null rows before making any billable requests. + (0..number_rows) + .map(|row| { + let values: Option> = 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 { + 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 Default for AiFunction { + 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 fmt::Debug for AiFunction { + 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 ScalarUDFImpl for AiFunction { + fn name(&self) -> &str { + Q::NAME + } + + fn signature(&self) -> &Signature { + &self.signature + } + + fn return_type(&self, _arg_types: &[DataType]) -> Result { + Ok(Q::return_type()) + } + + fn invoke_with_args(&self, _args: ScalarFunctionArgs) -> Result { + not_impl_err!("{} can only be called from async contexts", Q::NAME) + } +} + +#[async_trait] +impl AsyncScalarUDFImpl for AiFunction { + async fn invoke_async_with_args(&self, args: ScalarFunctionArgs) -> Result { + 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 = 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; + fn parse_answer(answer: &Value, criteria: &Self::Criteria) -> Result; +} + +#[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 { + 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; + + 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 { + 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 { + 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; + + 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 { + 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 { + 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> { + 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::>>()?; + // Allow small rounding differences without renormalizing the provider's distribution. + if (probabilities.iter().sum::() - 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 { + 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}]") + }) +} diff --git a/src/common/function/src/scalars/ai/tests.rs b/src/common/function/src/scalars/ai/tests.rs new file mode 100644 index 00000000000..0b001e22e7a --- /dev/null +++ b/src/common/function/src/scalars/ai/tests.rs @@ -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(&self) -> AiFunction { + 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(function: AiFunction) -> 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(criteria: &str) { + let function = AiFunction::::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::(r#"{"billing":null,"technical":"errors"}"#); + assert_shared_criteria::(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| 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::()); + 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:: { + 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:: { + 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::()); + 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| 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::()); + 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| 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 = (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::()); + 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| { + 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::()); + 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 = 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::()); + 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:: { + enabled: true, + api_key: None, + ..Default::default() + }); + let score_ctx = context(AiFunction:: { + enabled: true, + api_key: None, + ..Default::default() + }); + let too_many_choices: Map = (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::()); + ctx.register_udf( + AsyncScalarUDF::new(Arc::new(server.function::())).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); +} diff --git a/src/common/function/src/scalars/jev.rs b/src/common/function/src/scalars/jev.rs deleted file mode 100644 index 1c6475e6fac..00000000000 --- a/src/common/function/src/scalars/jev.rs +++ /dev/null @@ -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, - 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 { - 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 { - 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 { - Ok(DataType::Boolean) - } - - fn invoke_with_args(&self, _args: ScalarFunctionArgs) -> Result { - 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 { - let arrays = args - .args - .into_iter() - .map(|arg| arg.into_array(args.number_rows)) - .collect::>>()?; - 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> = stream::iter(requests).buffered(8).try_collect().await?; - Ok(ColumnarValue::Array(Arc::new(BooleanArray::from(matches)))) - } -} diff --git a/src/common/function/src/scalars/jev/tests.rs b/src/common/function/src/scalars/jev/tests.rs deleted file mode 100644 index 82d7ab03013..00000000000 --- a/src/common/function/src/scalars/jev/tests.rs +++ /dev/null @@ -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| 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); -} diff --git a/tests/cases/standalone/common/function/ai_functions.result b/tests/cases/standalone/common/function/ai_functions.result new file mode 100644 index 00000000000..057c9ecd990 --- /dev/null +++ b/tests/cases/standalone/common/function/ai_functions.result @@ -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) + diff --git a/tests/cases/standalone/common/function/ai_functions.sql b/tests/cases/standalone/common/function/ai_functions.sql new file mode 100644 index 00000000000..fd3e80d1e18 --- /dev/null +++ b/tests/cases/standalone/common/function/ai_functions.sql @@ -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'); diff --git a/tests/cases/standalone/common/function/jev.result b/tests/cases/standalone/common/function/jev.result deleted file mode 100644 index 4f6d91d2867..00000000000 --- a/tests/cases/standalone/common/function/jev.result +++ /dev/null @@ -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] - diff --git a/tests/cases/standalone/common/function/jev.sql b/tests/cases/standalone/common/function/jev.sql deleted file mode 100644 index f4c252c9273..00000000000 --- a/tests/cases/standalone/common/function/jev.sql +++ /dev/null @@ -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);