Skip to main content

common_function/scalars/
ai.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15#[cfg(test)]
16mod tests;
17
18use std::fmt;
19use std::hash::Hash;
20use std::marker::PhantomData;
21use std::sync::Arc;
22use std::time::Duration;
23
24use arrow::array::{Array, StringViewArray, new_null_array};
25use arrow::datatypes::DataType;
26use async_trait::async_trait;
27use datafusion_common::cast::as_string_view_array;
28use datafusion_common::{Result, ScalarValue, exec_datafusion_err, exec_err, not_impl_err};
29use datafusion_expr::async_udf::{AsyncScalarUDF, AsyncScalarUDFImpl};
30use datafusion_expr::{ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl, Signature, Volatility};
31use futures::{StreamExt, TryStreamExt, stream};
32use reqwest::Client;
33use serde::Serialize;
34use serde_json::{Map, Value, json};
35
36use crate::function_factory::ScalarFunctionFactory;
37use crate::function_registry::FunctionRegistry;
38
39/// Registers the experimental AI matching, classification, and rating functions.
40pub(crate) fn register(registry: &FunctionRegistry) {
41    AiFunction::<Noul>::register(registry);
42    AiFunction::<Choice>::register(registry);
43    AiFunction::<Score>::register(registry);
44}
45
46/// Shared asynchronous execution for AI functions using the Jev backend.
47#[derive(PartialEq, Eq, Hash)]
48struct AiFunction<Q> {
49    signature: Signature,
50    question: PhantomData<Q>,
51    enabled: bool,
52    api_key: Option<String>,
53    endpoint: String,
54    model: String,
55}
56
57struct AiRequest<'a, Q: AiQuestion> {
58    text: &'a str,
59    prompt: &'a str,
60    criteria: Arc<Q::Criteria>,
61}
62
63/// Borrows string arguments without expanding constants to the batch size.
64enum StringArgument<'a> {
65    Scalar(Option<&'a str>),
66    Array(&'a StringViewArray),
67}
68
69impl<'a> StringArgument<'a> {
70    fn try_new(arg: &'a ColumnarValue) -> Result<Self> {
71        match arg {
72            ColumnarValue::Scalar(ScalarValue::Utf8View(value)) => {
73                Ok(Self::Scalar(value.as_deref()))
74            }
75            ColumnarValue::Array(array) => Ok(Self::Array(as_string_view_array(array)?)),
76            _ => exec_err!("AI functions expect Utf8View arguments"),
77        }
78    }
79
80    fn value(&self, row: usize) -> Option<&'a str> {
81        match self {
82            Self::Scalar(value) => *value,
83            Self::Array(array) => {
84                let array = *array;
85                array.is_valid(row).then(|| array.value(row))
86            }
87        }
88    }
89}
90
91impl<Q: AiQuestion> AiFunction<Q> {
92    fn register(registry: &FunctionRegistry) {
93        registry.register(ScalarFunctionFactory {
94            name: Q::NAME.to_string(),
95            factory: Arc::new(|_| AsyncScalarUDF::new(Arc::new(Self::default())).into_scalar_udf()),
96        });
97    }
98
99    async fn evaluate(&self, client: &Client, request: AiRequest<'_, Q>) -> Result<ScalarValue> {
100        let mut question = json!({ "type": Q::TYPE, "instructions": request.prompt });
101        if Q::ARG_COUNT == 3 {
102            question["criteria"] = json!(request.criteria.as_ref());
103        }
104        let response: Value = client
105            .post(&self.endpoint)
106            .json(&json!({
107                "model": self.model,
108                "state": request.text,
109                "questions": { "matches": question }
110            }))
111            .send()
112            .await
113            .and_then(|response| response.error_for_status())
114            .map_err(|e| exec_datafusion_err!("{} request failed: {e}", Q::NAME))?
115            .json()
116            .await
117            .map_err(|e| exec_datafusion_err!("{} response is not valid JSON: {e}", Q::NAME))?;
118
119        let answer = &response["answers"]["matches"];
120        if answer["type"].as_str() != Some(Q::TYPE) {
121            return exec_err!(
122                "{} response must contain answers.matches with type {}",
123                Q::NAME,
124                Q::TYPE
125            );
126        }
127        Q::parse_answer(answer, request.criteria.as_ref())
128    }
129
130    fn prepare_requests<'a>(
131        &self,
132        args: &'a [ColumnarValue],
133        number_rows: usize,
134    ) -> Result<Vec<Option<AiRequest<'a, Q>>>> {
135        if args.len() != Q::ARG_COUNT {
136            return exec_err!("{} requires {} arguments", Q::NAME, Q::ARG_COUNT);
137        }
138        let args = args
139            .iter()
140            .map(StringArgument::try_new)
141            .collect::<Result<Vec<_>>>()?;
142        let scalar_criteria = args[2..]
143            .iter()
144            .all(|arg| matches!(arg, StringArgument::Scalar(_)));
145        let mut shared_criteria: Option<Arc<Q::Criteria>> = None;
146
147        // Validate all non-null rows before making any billable requests.
148        (0..number_rows)
149            .map(|row| {
150                let values: Option<Vec<_>> = args.iter().map(|arg| arg.value(row)).collect();
151                values
152                    .map(|values| {
153                        // Initialize lazily so NULL rows never validate otherwise invalid criteria.
154                        let criteria = match &shared_criteria {
155                            Some(criteria) => Arc::clone(criteria),
156                            None => {
157                                let criteria = Arc::new(Q::parse_criteria(&values[2..])?);
158                                if scalar_criteria {
159                                    shared_criteria = Some(Arc::clone(&criteria));
160                                }
161                                criteria
162                            }
163                        };
164                        Ok(AiRequest {
165                            text: values[0],
166                            prompt: values[1],
167                            criteria,
168                        })
169                    })
170                    .transpose()
171            })
172            .collect()
173    }
174
175    fn client(&self) -> Result<Client> {
176        if !self.enabled {
177            return exec_err!(
178                "{} is experimental; set GREPTIMEDB_EXPERIMENTAL_JEV=true to enable it",
179                Q::NAME
180            );
181        }
182        let key = self
183            .api_key
184            .as_deref()
185            .filter(|key| !key.trim().is_empty())
186            .ok_or_else(|| {
187                exec_datafusion_err!("{} requires the JEV_API_KEY environment variable", Q::NAME)
188            })?;
189        let mut authorization = reqwest::header::HeaderValue::from_str(&format!("Bearer {key}"))
190            .map_err(|_| exec_datafusion_err!("JEV_API_KEY is not a valid HTTP header value"))?;
191        authorization.set_sensitive(true);
192        let mut headers = reqwest::header::HeaderMap::new();
193        headers.insert(reqwest::header::AUTHORIZATION, authorization);
194        Client::builder()
195            .default_headers(headers)
196            .timeout(Duration::from_secs(30))
197            .build()
198            .map_err(|e| exec_datafusion_err!("failed to create {} HTTP client: {e}", Q::NAME))
199    }
200}
201
202impl<Q: AiQuestion> Default for AiFunction<Q> {
203    fn default() -> Self {
204        Self {
205            // External model evaluations must not be constant-folded during planning.
206            signature: Signature::exact(
207                vec![DataType::Utf8View; Q::ARG_COUNT],
208                Volatility::Volatile,
209            ),
210            question: PhantomData,
211            enabled: std::env::var("GREPTIMEDB_EXPERIMENTAL_JEV").as_deref() == Ok("true"),
212            api_key: std::env::var("JEV_API_KEY").ok(),
213            endpoint: std::env::var("JEV_ENDPOINT")
214                .unwrap_or_else(|_| "https://api.typesafe.ai/v1/systemone".to_string()),
215            model: std::env::var("JEV_MODEL").unwrap_or_else(|_| "jev-latest".to_string()),
216        }
217    }
218}
219
220impl<Q: AiQuestion> fmt::Debug for AiFunction<Q> {
221    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
222        // Credentials must never appear in plans or diagnostic output.
223        f.debug_struct("AiFunction")
224            .field("name", &Q::NAME)
225            .field("signature", &self.signature)
226            .field("enabled", &self.enabled)
227            .field("model", &self.model)
228            .finish_non_exhaustive()
229    }
230}
231
232impl<Q: AiQuestion> ScalarUDFImpl for AiFunction<Q> {
233    fn name(&self) -> &str {
234        Q::NAME
235    }
236
237    fn signature(&self) -> &Signature {
238        &self.signature
239    }
240
241    fn return_type(&self, _arg_types: &[DataType]) -> Result<DataType> {
242        Ok(Q::return_type())
243    }
244
245    fn invoke_with_args(&self, _args: ScalarFunctionArgs) -> Result<ColumnarValue> {
246        not_impl_err!("{} can only be called from async contexts", Q::NAME)
247    }
248}
249
250#[async_trait]
251impl<Q: AiQuestion> AsyncScalarUDFImpl for AiFunction<Q> {
252    async fn invoke_async_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
253        let rows = self.prepare_requests(&args.args, args.number_rows)?;
254
255        if rows.iter().all(Option::is_none) {
256            return Ok(ColumnarValue::Array(new_null_array(
257                &Q::return_type(),
258                args.number_rows,
259            )));
260        }
261
262        // Reuse connections within the batch; bound in-flight requests and preserve row order.
263        let client = self.client()?;
264        // Collect futures to avoid async_trait's higher-ranked lifetime inference issue.
265        let requests: Vec<_> = rows
266            .into_iter()
267            .map(|row| {
268                let client = &client;
269                async move {
270                    match row {
271                        Some(request) => self.evaluate(client, request).await,
272                        None => ScalarValue::try_from(&Q::return_type()),
273                    }
274                }
275            })
276            .collect();
277        // This limit is per expression/batch invocation, not per query or process.
278        // Concurrent partitions and queries can each have their own in-flight requests.
279        let answers: Vec<ScalarValue> = stream::iter(requests).buffered(8).try_collect().await?;
280        Ok(ColumnarValue::Array(ScalarValue::iter_to_array(answers)?))
281    }
282}
283
284/// The request criteria and scalar answer contract for an AI question type.
285trait AiQuestion: fmt::Debug + Eq + Hash + Send + Sync + 'static {
286    type Criteria: Serialize + Send + Sync;
287
288    const NAME: &'static str;
289    const TYPE: &'static str;
290    const ARG_COUNT: usize;
291
292    fn return_type() -> DataType;
293    fn parse_criteria(args: &[&str]) -> Result<Self::Criteria>;
294    fn parse_answer(answer: &Value, criteria: &Self::Criteria) -> Result<ScalarValue>;
295}
296
297#[derive(Debug, PartialEq, Eq, Hash)]
298struct Noul;
299
300impl AiQuestion for Noul {
301    type Criteria = ();
302
303    const NAME: &'static str = "ai_match";
304    const TYPE: &'static str = "noul";
305    const ARG_COUNT: usize = 2;
306
307    fn return_type() -> DataType {
308        DataType::Float64
309    }
310
311    fn parse_criteria(_args: &[&str]) -> Result<()> {
312        Ok(())
313    }
314
315    fn parse_answer(answer: &Value, _criteria: &()) -> Result<ScalarValue> {
316        numeric_answer(Self::NAME, &answer["noul"], "noul", 1.0)
317            .map(|probability| ScalarValue::Float64(Some(probability)))
318    }
319}
320
321#[derive(Debug, PartialEq, Eq, Hash)]
322struct Choice;
323
324impl AiQuestion for Choice {
325    type Criteria = Map<String, Value>;
326
327    const NAME: &'static str = "ai_choose";
328    const TYPE: &'static str = "choice";
329    const ARG_COUNT: usize = 3;
330
331    fn return_type() -> DataType {
332        DataType::Utf8
333    }
334
335    fn parse_criteria(args: &[&str]) -> Result<Self::Criteria> {
336        let criteria: Value = serde_json::from_str(args[0])
337            .map_err(|e| exec_datafusion_err!("ai_choose criteria is not valid JSON: {e}"))?;
338        match criteria {
339            Value::Object(options)
340                if (1..=255).contains(&options.len())
341                    && options.values().all(|v| v.is_null() || is_description(v)) =>
342            {
343                Ok(options)
344            }
345            _ => exec_err!(
346                "ai_choose criteria must be a JSON object with 1 to 255 options; descriptions must be strings, objects, arrays, or null"
347            ),
348        }
349    }
350
351    fn parse_answer(answer: &Value, criteria: &Self::Criteria) -> Result<ScalarValue> {
352        let choice = answer["choice"]
353            .as_str()
354            .filter(|choice| criteria.contains_key(*choice))
355            .ok_or_else(|| {
356                exec_datafusion_err!("ai_choose response must contain a choice from the criteria")
357            })?;
358        Ok(ScalarValue::Utf8(Some(choice.to_string())))
359    }
360}
361
362#[derive(Debug, PartialEq, Eq, Hash)]
363struct Score;
364
365impl AiQuestion for Score {
366    type Criteria = Vec<Value>;
367
368    const NAME: &'static str = "ai_score";
369    const TYPE: &'static str = "score";
370    const ARG_COUNT: usize = 3;
371
372    fn return_type() -> DataType {
373        DataType::BinaryView
374    }
375
376    fn parse_criteria(args: &[&str]) -> Result<Self::Criteria> {
377        let criteria: Value = serde_json::from_str(args[0])
378            .map_err(|e| exec_datafusion_err!("ai_score criteria is not valid JSON: {e}"))?;
379        match criteria {
380            Value::Array(levels)
381                if (2..=10).contains(&levels.len()) && levels.iter().all(is_description) =>
382            {
383                Ok(levels)
384            }
385            _ => exec_err!(
386                "ai_score criteria must be a JSON array with 2 to 10 levels; descriptions must be strings, objects, or arrays"
387            ),
388        }
389    }
390
391    fn parse_answer(answer: &Value, criteria: &Self::Criteria) -> Result<ScalarValue> {
392        let score = numeric_answer(
393            Self::NAME,
394            &answer["score"],
395            "score",
396            (criteria.len() - 1) as f64,
397        )?;
398        let confidence = numeric_answer(Self::NAME, &answer["confidence"], "confidence", 1.0)?;
399        let probabilities = score_probabilities(&answer["probabilities"], criteria.len())?;
400        let object = jsonb::Object::from([
401            ("score".to_string(), jsonb::Value::from(score)),
402            ("confidence".to_string(), jsonb::Value::from(confidence)),
403            (
404                "probabilities".to_string(),
405                jsonb::Value::Array(probabilities.into_iter().map(jsonb::Value::from).collect()),
406            ),
407        ]);
408        Ok(ScalarValue::BinaryView(Some(
409            jsonb::Value::Object(object).to_vec(),
410        )))
411    }
412}
413
414fn score_probabilities(distribution: &Value, level_count: usize) -> Result<Vec<f64>> {
415    if distribution
416        .as_object()
417        .is_none_or(|probabilities| probabilities.len() != level_count)
418    {
419        return exec_err!(
420            "ai_score response must contain probabilities for all {level_count} levels"
421        );
422    }
423    let probabilities = (0..level_count)
424        .map(|level| {
425            numeric_answer(
426                Score::NAME,
427                &distribution[level.to_string()],
428                "probability",
429                1.0,
430            )
431        })
432        .collect::<Result<Vec<_>>>()?;
433    // Allow small rounding differences without renormalizing the provider's distribution.
434    if (probabilities.iter().sum::<f64>() - 1.0).abs() > 1e-6 {
435        return exec_err!("ai_score response probabilities must sum to 1 within 1e-6");
436    }
437    Ok(probabilities)
438}
439
440fn is_description(description: &Value) -> bool {
441    matches!(
442        description,
443        Value::String(_) | Value::Object(_) | Value::Array(_)
444    )
445}
446
447fn numeric_answer(name: &str, answer: &Value, field: &str, max: f64) -> Result<f64> {
448    answer
449        .as_f64()
450        .filter(|score| (0.0..=max).contains(score))
451        .ok_or_else(|| {
452            exec_datafusion_err!("{name} response must contain a finite {field} in [0, {max}]")
453        })
454}