1#[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
39pub(crate) fn register(registry: &FunctionRegistry) {
41 AiFunction::<Noul>::register(registry);
42 AiFunction::<Choice>::register(registry);
43 AiFunction::<Score>::register(registry);
44}
45
46#[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
63enum 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 (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 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 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 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 let client = self.client()?;
264 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 let answers: Vec<ScalarValue> = stream::iter(requests).buffered(8).try_collect().await?;
280 Ok(ColumnarValue::Array(ScalarValue::iter_to_array(answers)?))
281 }
282}
283
284trait 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 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}