This is an automated email from the ASF dual-hosted git repository.
yuxiqian pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/flink-cdc.git
The following commit(s) were added to refs/heads/master by this push:
new b11d6fdf63 [FLINK-40572][runtime] Support dynamic AI model selection
(#4525)
b11d6fdf63 is described below
commit b11d6fdf6346520e720396f25e5b5162ce5435b9
Author: haruki <[email protected]>
AuthorDate: Wed Sep 9 15:46:08 2026 +0800
[FLINK-40572][runtime] Support dynamic AI model selection (#4525)
---
docs/content.zh/docs/core-concept/ai-model.md | 12 +-
docs/content/docs/core-concept/ai-model.md | 12 +-
.../flink/translator/TransformTranslator.java | 49 -----
.../flink/FlinkPipelineAiFunctionITCase.java | 14 ++
.../cdc/runtime/functions/impl/AiFunctions.java | 149 ++++++++++----
.../transform/ProjectionColumnProcessor.java | 6 +-
.../transform/TransformExpressionCompiler.java | 23 +--
.../transform/TransformFilterProcessor.java | 9 +-
.../flink/cdc/runtime/parser/JaninoCompiler.java | 20 +-
.../flink/cdc/runtime/parser/TransformParser.java | 209 -------------------
.../runtime/functions/impl/AiFunctionsTest.java | 115 ++++++++---
.../cdc/runtime/parser/AiFunctionParserTest.java | 225 +++------------------
12 files changed, 282 insertions(+), 561 deletions(-)
diff --git a/docs/content.zh/docs/core-concept/ai-model.md
b/docs/content.zh/docs/core-concept/ai-model.md
index d809b3224a..70e016be91 100644
--- a/docs/content.zh/docs/core-concept/ai-model.md
+++ b/docs/content.zh/docs/core-concept/ai-model.md
@@ -28,7 +28,17 @@ AI 模型可用于 transform 表达式中的文本生成、文本分析、embedd
## AI Functions
-模型名称必须是字符串常量,并引用 `pipeline.model` 中声明的模型。文本、embedding 和图片函数分别要求模型客户端实现对应的
capability;Pipeline 会在执行前校验引用模型的 capability 是否匹配。
+模型参数可以是任意 `STRING` 表达式,并会针对每条记录求值。因此可以使用字段、`IF` 或 `CASE` 动态选择模型。仅在实际调用 AI
函数时,才会根据求值结果查找 `pipeline.model` 中声明的模型;如果选中的模型未声明,或没有实现函数所需的
capability,当前记录会在运行时报错。
+
+例如,下面的表达式会根据每条记录的优先级选择模型:
+
+```sql
+AI_COMPLETE(
+ IF(priority = 'high', 'powerful_model', 'economical_model'),
+ content,
+ '总结输入内容'
+)
+```
所有文本函数都会将模型返回的 JSON 解析为 `VARIANT`。
diff --git a/docs/content/docs/core-concept/ai-model.md
b/docs/content/docs/core-concept/ai-model.md
index 1ad5a739fe..7267098c50 100644
--- a/docs/content/docs/core-concept/ai-model.md
+++ b/docs/content/docs/core-concept/ai-model.md
@@ -29,7 +29,17 @@ image understanding.
## AI Functions
-The model name must be a string constant that refers to a model declared in
`pipeline.model`. Text functions require a model client that implements text
generation, while embedding and image functions require their corresponding
capabilities. The pipeline validates the referenced model capability before
execution.
+The model argument accepts any `STRING` expression and is evaluated for each
record. This enables dynamic model selection with a column, `IF`, or `CASE`.
The selected name is resolved against the models declared in `pipeline.model`
only when the AI function is invoked. If the selected model is undeclared or
does not provide the capability required by the function, processing of that
record fails at runtime.
+
+For example, the following expression chooses a model based on each record's
priority:
+
+```sql
+AI_COMPLETE(
+ IF(priority = 'high', 'powerful_model', 'economical_model'),
+ content,
+ 'Summarize the input'
+)
+```
All text functions return `VARIANT` values parsed from the model's JSON
response.
diff --git
a/flink-cdc-composer/src/main/java/org/apache/flink/cdc/composer/flink/translator/TransformTranslator.java
b/flink-cdc-composer/src/main/java/org/apache/flink/cdc/composer/flink/translator/TransformTranslator.java
index 5dc9b268f7..99be446f20 100644
---
a/flink-cdc-composer/src/main/java/org/apache/flink/cdc/composer/flink/translator/TransformTranslator.java
+++
b/flink-cdc-composer/src/main/java/org/apache/flink/cdc/composer/flink/translator/TransformTranslator.java
@@ -35,17 +35,14 @@ import
org.apache.flink.cdc.runtime.operators.transform.PostTransformOperator;
import
org.apache.flink.cdc.runtime.operators.transform.PostTransformOperatorBuilder;
import org.apache.flink.cdc.runtime.operators.transform.PreTransformOperator;
import
org.apache.flink.cdc.runtime.operators.transform.PreTransformOperatorBuilder;
-import org.apache.flink.cdc.runtime.parser.TransformParser;
import org.apache.flink.cdc.runtime.typeutils.EventTypeInfo;
import org.apache.flink.streaming.api.datastream.DataStream;
import org.apache.flink.streaming.api.environment.StreamExecutionEnvironment;
import java.util.Collections;
-import java.util.HashSet;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
-import java.util.Set;
import java.util.stream.Collectors;
/**
@@ -67,8 +64,6 @@ public class TransformTranslator {
if (transforms.isEmpty()) {
return input;
}
- validateModelReferences(
- transforms, models, getUserDefinedFunctionNames(udfFunctions,
models));
return input.transform(
"Transform:Schema",
new EventTypeInfo(),
@@ -147,8 +142,6 @@ public class TransformTranslator {
.map(this::modelToUDFTuple)
.collect(Collectors.toList()));
Map<String, AiModelClient> modelClients = loadModelClients(models,
env);
- validateModelCapabilities(
- transforms, modelClients,
getUserDefinedFunctionNames(udfFunctions, models));
postTransformFunctionBuilder.addModelClients(modelClients);
return input.transform(
"Transform:Data", new EventTypeInfo(),
postTransformFunctionBuilder.build())
@@ -191,48 +184,6 @@ public class TransformTranslator {
return clients;
}
- private void validateModelReferences(
- List<TransformDef> transforms,
- List<ModelDef> models,
- Set<String> userDefinedFunctionNames) {
- Set<String> clientModelNames =
- models.stream()
- .filter(model -> !model.isLegacy())
- .map(ModelDef::getName)
- .collect(Collectors.toSet());
- for (TransformDef transform : transforms) {
- TransformParser.validateAiModelReferences(
- transform.getProjection(),
- transform.getFilter(),
- clientModelNames,
- userDefinedFunctionNames);
- }
- }
-
- private void validateModelCapabilities(
- List<TransformDef> transforms,
- Map<String, AiModelClient> modelClients,
- Set<String> userDefinedFunctionNames) {
- for (TransformDef transform : transforms) {
- TransformParser.validateAiModelCapabilities(
- transform.getProjection(),
- transform.getFilter(),
- modelClients,
- userDefinedFunctionNames);
- }
- }
-
- private Set<String> getUserDefinedFunctionNames(
- List<UdfDef> udfFunctions, List<ModelDef> models) {
- Set<String> functionNames = new HashSet<>();
- udfFunctions.stream().map(UdfDef::getName).forEach(functionNames::add);
- models.stream()
- .filter(ModelDef::isLegacy)
- .map(ModelDef::getName)
- .forEach(functionNames::add);
- return functionNames;
- }
-
private Tuple3<String, String, Map<String, String>>
udfDefToUDFTuple(UdfDef udf) {
return Tuple3.of(udf.getName(), udf.getClasspath(), udf.getOptions());
}
diff --git
a/flink-cdc-composer/src/test/java/org/apache/flink/cdc/composer/flink/FlinkPipelineAiFunctionITCase.java
b/flink-cdc-composer/src/test/java/org/apache/flink/cdc/composer/flink/FlinkPipelineAiFunctionITCase.java
index d0610ed9f5..22d79749b7 100644
---
a/flink-cdc-composer/src/test/java/org/apache/flink/cdc/composer/flink/FlinkPipelineAiFunctionITCase.java
+++
b/flink-cdc-composer/src/test/java/org/apache/flink/cdc/composer/flink/FlinkPipelineAiFunctionITCase.java
@@ -128,6 +128,20 @@ class FlinkPipelineAiFunctionITCase {
"Dummy model closed.");
}
+ @Test
+ void testDynamicModelSelectionInProjection() throws Exception {
+ String[] output =
+ runAiFunctionTest(
+ "id, content, "
+ + "AI_COMPLETE(IF(id = 1, 'testModel',
'missingModel'), content, 'Complete the text') AS completed",
+ List.of(ModelDef.of("testModel", "dummy",
Collections.emptyMap())));
+
+ assertThat(output)
+ .containsExactly(
+
"CreateTableEvent{tableId=default_namespace.default_schema.mytable1,
schema=columns={`id` INT NOT NULL,`content` STRING,`completed` VARIANT},
primaryKeys=id, options=()}",
+
"DataChangeEvent{tableId=default_namespace.default_schema.mytable1, before=[],
after=[1, I love this product, {\"result\":\"dummy response\"}], op=INSERT,
meta=()}");
+ }
+
@Test
void testSpecializedTextAiFunctionsInProjection() throws Exception {
String[] output =
diff --git
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/functions/impl/AiFunctions.java
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/functions/impl/AiFunctions.java
index 2c21ba1c71..2d874ac606 100644
---
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/functions/impl/AiFunctions.java
+++
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/functions/impl/AiFunctions.java
@@ -31,6 +31,7 @@ import
org.apache.flink.shaded.guava31.com.google.common.primitives.Floats;
import java.io.IOException;
import java.util.List;
+import java.util.Map;
/** Built-in AI functions available to transform expressions. */
public class AiFunctions {
@@ -39,53 +40,89 @@ public class AiFunctions {
private AiFunctions() {}
- public static BinaryVariant aiComplete(AiModelClient model, String input,
String systemPrompt) {
- return generateText(model, AiTextFunctionDef.AI_COMPLETE, input,
systemPrompt);
+ public static BinaryVariant aiComplete(
+ String modelName,
+ String input,
+ String systemPrompt,
+ Map<String, AiModelClient> modelClients) {
+ return generateText(
+ modelClients, modelName, AiTextFunctionDef.AI_COMPLETE, input,
systemPrompt);
}
- public static BinaryVariant aiClassify(AiModelClient model, String input,
String labels) {
- return generateText(model, AiTextFunctionDef.AI_CLASSIFY, input,
labels);
+ public static BinaryVariant aiClassify(
+ String modelName,
+ String input,
+ String labels,
+ Map<String, AiModelClient> modelClients) {
+ return generateText(modelClients, modelName,
AiTextFunctionDef.AI_CLASSIFY, input, labels);
}
public static BinaryVariant aiTranslate(
- AiModelClient model, String input, String sourceLang, String
targetLang) {
- return generateText(model, AiTextFunctionDef.AI_TRANSLATE, input,
sourceLang, targetLang);
+ String modelName,
+ String input,
+ String sourceLang,
+ String targetLang,
+ Map<String, AiModelClient> modelClients) {
+ return generateText(
+ modelClients,
+ modelName,
+ AiTextFunctionDef.AI_TRANSLATE,
+ input,
+ sourceLang,
+ targetLang);
}
- public static BinaryVariant aiSummarize(AiModelClient model, String input,
int maxLength) {
- return generateText(model, AiTextFunctionDef.AI_SUMMARIZE, input,
maxLength);
+ public static BinaryVariant aiSummarize(
+ String modelName,
+ String input,
+ int maxLength,
+ Map<String, AiModelClient> modelClients) {
+ return generateText(
+ modelClients, modelName, AiTextFunctionDef.AI_SUMMARIZE,
input, maxLength);
}
- public static BinaryVariant aiSentiment(AiModelClient model, String input)
{
- return generateText(model, AiTextFunctionDef.AI_SENTIMENT, input);
+ public static BinaryVariant aiSentiment(
+ String modelName, String input, Map<String, AiModelClient>
modelClients) {
+ return generateText(modelClients, modelName,
AiTextFunctionDef.AI_SENTIMENT, input);
}
- public static BinaryVariant aiExtract(AiModelClient model, String input,
String schema) {
- return generateText(model, AiTextFunctionDef.AI_EXTRACT, input,
schema);
+ public static BinaryVariant aiExtract(
+ String modelName,
+ String input,
+ String schema,
+ Map<String, AiModelClient> modelClients) {
+ return generateText(modelClients, modelName,
AiTextFunctionDef.AI_EXTRACT, input, schema);
}
- public static BinaryVariant aiMask(AiModelClient model, String input,
String entities) {
- return generateText(model, AiTextFunctionDef.AI_MASK, input, entities);
+ public static BinaryVariant aiMask(
+ String modelName,
+ String input,
+ String entities,
+ Map<String, AiModelClient> modelClients) {
+ return generateText(modelClients, modelName,
AiTextFunctionDef.AI_MASK, input, entities);
}
private static BinaryVariant generateText(
- AiModelClient model,
+ Map<String, AiModelClient> modelClients,
+ String modelName,
AiTextFunctionDef function,
String input,
Object... promptArguments) {
if (input == null) {
return null;
}
- if (!(model instanceof SupportsTextGeneration)) {
- throw new UnsupportedOperationException(
- "Model " + model.getClass().getName() + " does not support
text generation");
- }
+ SupportsTextGeneration model =
+ resolveModel(
+ modelClients,
+ modelName,
+ function.getFunctionName(),
+ SupportsTextGeneration.class);
String prompt =
function.buildPrompt(promptArguments)
+ "\n"
+ buildOutputSchemaHint(function.getOutputType());
- String json = ((SupportsTextGeneration) model).generate(prompt, input);
+ String json = model.generate(prompt, input);
if (json == null) {
return null;
}
@@ -101,43 +138,77 @@ public class AiFunctions {
}
}
- public static List<Float> aiEmbed(AiModelClient model, String input) {
+ public static List<Float> aiEmbed(
+ String modelName, String input, Map<String, AiModelClient>
modelClients) {
if (input == null) {
return null;
}
- if (!(model instanceof SupportsEmbedding)) {
- throw new UnsupportedOperationException(
- "Model " + model.getClass().getName() + " does not support
embedding");
- }
- float[] embedding = ((SupportsEmbedding) model).embed(input);
+ SupportsEmbedding model =
+ resolveModel(modelClients, modelName, "AI_EMBED",
SupportsEmbedding.class);
+ float[] embedding = model.embed(input);
return embedding == null ? null : Floats.asList(embedding);
}
/** Dispatches image-to-text AI functions. */
- public static String aiImageComplete(AiModelClient model, byte[] image,
String prompt) {
+ public static String aiImageComplete(
+ String modelName,
+ byte[] image,
+ String prompt,
+ Map<String, AiModelClient> modelClients) {
if (image == null) {
return null;
}
- if (!(model instanceof SupportsImageTextGeneration)) {
- throw new UnsupportedOperationException(
- "Model "
- + model.getClass().getName()
- + " does not support image text generation");
- }
- return ((SupportsImageTextGeneration)
model).generateTextFromImage(image, prompt);
+ SupportsImageTextGeneration model =
+ resolveModel(
+ modelClients,
+ modelName,
+ "AI_IMAGE_COMPLETE",
+ SupportsImageTextGeneration.class);
+ return model.generateTextFromImage(image, prompt);
}
/** Dispatches image embedding AI functions. */
- public static List<Float> aiImageEmbed(AiModelClient model, byte[] image) {
+ public static List<Float> aiImageEmbed(
+ String modelName, byte[] image, Map<String, AiModelClient>
modelClients) {
if (image == null) {
return null;
}
- if (!(model instanceof SupportsImageEmbedding)) {
+ SupportsImageEmbedding model =
+ resolveModel(
+ modelClients, modelName, "AI_IMAGE_EMBED",
SupportsImageEmbedding.class);
+ float[] embedding = model.embedImage(image);
+ return embedding == null ? null : Floats.asList(embedding);
+ }
+
+ private static <T> T resolveModel(
+ Map<String, AiModelClient> modelClients,
+ String modelName,
+ String functionName,
+ Class<T> requiredCapability) {
+ if (modelName == null) {
+ throw new IllegalArgumentException(
+ "Model name referenced by " + functionName + " must not be
null.");
+ }
+ AiModelClient model = modelClients.get(modelName);
+ if (model == null) {
+ throw new IllegalArgumentException(
+ "Model '"
+ + modelName
+ + "' referenced by "
+ + functionName
+ + " has not been declared.");
+ }
+ if (!requiredCapability.isInstance(model)) {
throw new UnsupportedOperationException(
- "Model " + model.getClass().getName() + " does not support
image embedding");
+ "Model '"
+ + modelName
+ + "' could not be used in "
+ + functionName
+ + " because it does not implement "
+ + requiredCapability.getSimpleName()
+ + " interface.");
}
- float[] embedding = ((SupportsImageEmbedding) model).embedImage(image);
- return embedding == null ? null : Floats.asList(embedding);
+ return requiredCapability.cast(model);
}
private static String truncateInvalidJsonResponse(String response) {
diff --git
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/operators/transform/ProjectionColumnProcessor.java
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/operators/transform/ProjectionColumnProcessor.java
index 2dc86b37be..d4697007ea 100644
---
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/operators/transform/ProjectionColumnProcessor.java
+++
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/operators/transform/ProjectionColumnProcessor.java
@@ -66,7 +66,7 @@ public class ProjectionColumnProcessor {
this.transformExpressionKey = generateTransformExpressionKey();
this.expressionEvaluator =
TransformExpressionCompiler.compileExpression(
- transformExpressionKey, udfDescriptors, modelClients);
+ transformExpressionKey, udfDescriptors);
this.udfFunctionInstances = udfFunctionInstances;
}
@@ -148,8 +148,8 @@ public class ProjectionColumnProcessor {
// 3 - Add UDF function instances
params.addAll(udfFunctionInstances);
- // 4 - Add AI model client instances
- params.addAll(modelClients.values());
+ // 4 - Add AI model clients
+ params.add(modelClients);
return params.toArray();
}
diff --git
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/operators/transform/TransformExpressionCompiler.java
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/operators/transform/TransformExpressionCompiler.java
index 47e42eb90b..33a70a9ce9 100644
---
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/operators/transform/TransformExpressionCompiler.java
+++
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/operators/transform/TransformExpressionCompiler.java
@@ -18,8 +18,8 @@
package org.apache.flink.cdc.runtime.operators.transform;
import org.apache.flink.api.common.InvalidProgramException;
-import org.apache.flink.cdc.common.model.AiModelClient;
import
org.apache.flink.cdc.runtime.operators.transform.exceptions.TransformException;
+import org.apache.flink.cdc.runtime.parser.JaninoCompiler;
import org.apache.flink.util.FlinkRuntimeException;
import org.apache.flink.shaded.guava31.com.google.common.cache.Cache;
@@ -31,7 +31,6 @@ import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import java.util.ArrayList;
-import java.util.Collections;
import java.util.List;
import java.util.Map;
@@ -57,20 +56,6 @@ public class TransformExpressionCompiler {
/** Compiles an expression code to a janino {@link ExpressionEvaluator}. */
public static ExpressionEvaluator compileExpression(
TransformExpressionKey key, List<UserDefinedFunctionDescriptor>
udfDescriptors) {
- return compileExpression(key, udfDescriptors, Collections.emptyMap());
- }
-
- /**
- * Compiles an expression code to a janino {@link ExpressionEvaluator},
with additional {@link
- * AiModelClient} instances appended after UDF instances.
- *
- * <p>{@code modelClients} maps model names (e.g. {@code myModel}) to the
corresponding client
- * instances.
- */
- public static ExpressionEvaluator compileExpression(
- TransformExpressionKey key,
- List<UserDefinedFunctionDescriptor> udfDescriptors,
- Map<String, AiModelClient> modelClients) {
try {
return COMPILED_EXPRESSION_CACHE.get(
key,
@@ -85,10 +70,8 @@ public class TransformExpressionCompiler {
argumentClasses.add(Class.forName(udfFunction.getClasspath()));
}
- for (String paramName : modelClients.keySet()) {
- argumentNames.add(paramName);
- argumentClasses.add(AiModelClient.class);
- }
+
argumentNames.add(JaninoCompiler.DEFAULT_AI_MODEL_CLIENTS);
+ argumentClasses.add(Map.class);
// Input args
expressionEvaluator.setParameters(
diff --git
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/operators/transform/TransformFilterProcessor.java
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/operators/transform/TransformFilterProcessor.java
index 02b2c51716..6e4eac45f3 100644
---
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/operators/transform/TransformFilterProcessor.java
+++
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/operators/transform/TransformFilterProcessor.java
@@ -30,6 +30,7 @@ import org.codehaus.janino.ExpressionEvaluator;
import java.lang.reflect.InvocationTargetException;
import java.util.ArrayList;
+import java.util.Collections;
import java.util.HashMap;
import java.util.HashSet;
import java.util.LinkedHashSet;
@@ -72,7 +73,7 @@ public class TransformFilterProcessor {
this.decimalPrecisionMode = decimalPrecisionMode;
this.udfFunctionInstances = udfFunctionInstances;
this.supportedMetadataColumns = supportedMetadataColumns;
- this.modelClients = modelClients;
+ this.modelClients = modelClients == null ? Collections.emptyMap() :
modelClients;
if (isNoOp) {
this.transformExpressionKey = null;
@@ -87,7 +88,7 @@ public class TransformFilterProcessor {
.toArray(new SupportedMetadataColumn[0]));
this.expressionEvaluator =
TransformExpressionCompiler.compileExpression(
- transformExpressionKey, udfDescriptors,
modelClients);
+ transformExpressionKey, udfDescriptors);
}
}
@@ -223,8 +224,8 @@ public class TransformFilterProcessor {
// 3 - Add UDF function instances
params.addAll(udfFunctionInstances);
- // 4 - Add AI model client instances
- params.addAll(modelClients.values());
+ // 4 - Add AI model clients
+ params.add(modelClients);
return params.toArray();
}
diff --git
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/parser/JaninoCompiler.java
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/parser/JaninoCompiler.java
index 627cde88fa..340e7e61eb 100644
---
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/parser/JaninoCompiler.java
+++
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/parser/JaninoCompiler.java
@@ -119,6 +119,7 @@ public class JaninoCompiler {
public static final String DEFAULT_EPOCH_TIME = "__epoch_time__";
public static final String DEFAULT_TIME_ZONE = "__time_zone__";
+ public static final String DEFAULT_AI_MODEL_CLIENTS =
"__ai_model_clients__";
private static final String[] BUILTIN_FUNCTION_MODULES = {
"Ai", "Arithmetic", "Casting", "Comparison", "Logical", "String",
"Struct", "Temporal"
@@ -264,6 +265,10 @@ public class JaninoCompiler {
} else if (TIMEZONE_REQUIRED_TEMPORAL_CONVERSION_FUNCTIONS.contains(
sqlBasicCall.getOperator().getName().toUpperCase())) {
atoms.add(new Java.AmbiguousName(Location.NOWHERE, new String[]
{DEFAULT_TIME_ZONE}));
+ } else if (isAiFunction(functionName)) {
+ atoms.add(
+ new Java.AmbiguousName(
+ Location.NOWHERE, new String[]
{DEFAULT_AI_MODEL_CLIENTS}));
}
return sqlBasicCallToJaninoRvalue(context, sqlBasicCall,
atoms.toArray(new Java.Rvalue[0]));
}
@@ -1002,13 +1007,6 @@ public class JaninoCompiler {
return castExpressionToInferredType(
context, sqlBasicCall,
generateFunctionOperation("element", atoms));
} else {
- if (isAiFunction(operationName) && atoms.length >= 1) {
- if (!(sqlBasicCall.operand(0) instanceof
SqlCharStringLiteral)) {
- throw new ParseException(
- "The model argument of an AI function must be a
string constant.");
- }
- rewriteAiFunctionModelArg(atoms);
- }
return new Java.MethodInvocation(
Location.NOWHERE,
null,
@@ -1122,14 +1120,6 @@ public class JaninoCompiler {
return false;
}
- private static void rewriteAiFunctionModelArg(Java.Rvalue[] atoms) {
- String modelName = atoms[0].toString();
- if (modelName.startsWith("\"") && modelName.endsWith("\"")) {
- modelName = modelName.substring(1, modelName.length() - 1);
- }
- atoms[0] = new Java.AmbiguousName(Location.NOWHERE, new String[]
{modelName});
- }
-
private static Java.Rvalue generateTimezoneFreeTemporalFunctionOperation(
Context context, String operationName) {
return new Java.MethodInvocation(
diff --git
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/parser/TransformParser.java
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/parser/TransformParser.java
index 5993691af9..ebddd1de8a 100644
---
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/parser/TransformParser.java
+++
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/parser/TransformParser.java
@@ -18,19 +18,11 @@
package org.apache.flink.cdc.runtime.parser;
import org.apache.flink.api.common.io.ParseException;
-import org.apache.flink.cdc.common.model.AiModelClient;
-import org.apache.flink.cdc.common.model.abilities.SupportsEmbedding;
-import org.apache.flink.cdc.common.model.abilities.SupportsImageEmbedding;
-import org.apache.flink.cdc.common.model.abilities.SupportsImageTextGeneration;
-import org.apache.flink.cdc.common.model.abilities.SupportsTextGeneration;
import org.apache.flink.cdc.common.pipeline.DecimalPrecisionMode;
import org.apache.flink.cdc.common.schema.Column;
import org.apache.flink.cdc.common.source.SupportedMetadataColumn;
import org.apache.flink.cdc.common.types.DataType;
import org.apache.flink.cdc.common.utils.Preconditions;
-import org.apache.flink.cdc.runtime.ai.AiEmbeddingFunctionDef;
-import org.apache.flink.cdc.runtime.ai.AiImageFunctionDef;
-import org.apache.flink.cdc.runtime.ai.AiTextFunctionDef;
import org.apache.flink.cdc.runtime.operators.transform.ProjectionColumn;
import
org.apache.flink.cdc.runtime.operators.transform.UserDefinedFunctionDescriptor;
import org.apache.flink.cdc.runtime.parser.metadata.AiFunctionSqlOperatorTable;
@@ -56,7 +48,6 @@ import org.apache.calcite.schema.SchemaPlus;
import org.apache.calcite.schema.impl.ScalarFunctionImpl;
import org.apache.calcite.sql.SqlBasicCall;
import org.apache.calcite.sql.SqlCall;
-import org.apache.calcite.sql.SqlCharStringLiteral;
import org.apache.calcite.sql.SqlFunction;
import org.apache.calcite.sql.SqlFunctionCategory;
import org.apache.calcite.sql.SqlIdentifier;
@@ -906,206 +897,6 @@ public class TransformParser {
return parseSelect(statement.toString());
}
- /** Validates model arguments and references in the supported AI
functions. */
- public static void validateAiModelReferences(
- @Nullable String projection, @Nullable String filter, Set<String>
declaredModelNames) {
- validateAiModelReferences(projection, filter, declaredModelNames,
Collections.emptySet());
- }
-
- /** Validates model arguments and references in AI functions not shadowed
by a UDF. */
- public static void validateAiModelReferences(
- @Nullable String projection,
- @Nullable String filter,
- Set<String> declaredModelNames,
- Set<String> userDefinedFunctionNames) {
- if (!isNullOrWhitespaceOnly(projection)) {
- validateAiModelReferences(
- parseProjectionExpression(projection),
- declaredModelNames,
- userDefinedFunctionNames);
- }
- if (!isNullOrWhitespaceOnly(filter)) {
- validateAiModelReferences(
- parseFilterExpression(filter), declaredModelNames,
userDefinedFunctionNames);
- }
- }
-
- /** Validates that referenced models provide the capability required by
each AI function. */
- public static void validateAiModelCapabilities(
- @Nullable String projection,
- @Nullable String filter,
- Map<String, AiModelClient> modelClients) {
- validateAiModelCapabilities(projection, filter, modelClients,
Collections.emptySet());
- }
-
- /** Validates model capabilities for AI functions not shadowed by a UDF. */
- public static void validateAiModelCapabilities(
- @Nullable String projection,
- @Nullable String filter,
- Map<String, AiModelClient> modelClients,
- Set<String> userDefinedFunctionNames) {
- if (!isNullOrWhitespaceOnly(projection)) {
- validateAiModelCapabilities(
- parseProjectionExpression(projection), modelClients,
userDefinedFunctionNames);
- }
- if (!isNullOrWhitespaceOnly(filter)) {
- validateAiModelCapabilities(
- parseFilterExpression(filter), modelClients,
userDefinedFunctionNames);
- }
- }
-
- private static void validateAiModelReferences(
- SqlNode node, Set<String> declaredModelNames, Set<String>
userDefinedFunctionNames) {
- if (node instanceof SqlCall) {
- SqlCall call = (SqlCall) node;
- String functionName = call.getOperator().getName();
- if (isAiFunction(functionName)
- && !isUserDefinedFunction(functionName,
userDefinedFunctionNames)) {
- if (call.operandCount() == 0) {
- return;
- }
- String modelName = resolveAiModelName(call);
- Preconditions.checkArgument(
- declaredModelNames.contains(modelName),
- "Model '%s' referenced by %s has not been declared.",
- modelName,
- call.getOperator().getName());
- }
- for (SqlNode operand : call.getOperandList()) {
- if (operand != null) {
- validateAiModelReferences(
- operand, declaredModelNames,
userDefinedFunctionNames);
- }
- }
- } else if (node instanceof SqlNodeList) {
- for (SqlNode child : (SqlNodeList) node) {
- validateAiModelReferences(child, declaredModelNames,
userDefinedFunctionNames);
- }
- }
- }
-
- private static void validateAiModelCapabilities(
- SqlNode node,
- Map<String, AiModelClient> modelClients,
- Set<String> userDefinedFunctionNames) {
- if (node instanceof SqlCall) {
- SqlCall call = (SqlCall) node;
- String functionName = call.getOperator().getName();
- if (isAiFunction(functionName)
- && !isUserDefinedFunction(functionName,
userDefinedFunctionNames)
- && call.operandCount() > 0) {
- String modelName = resolveAiModelName(call);
- AiModelClient modelClient = modelClients.get(modelName);
- Preconditions.checkArgument(
- modelClient != null,
- "Model '%s' referenced by %s has not been declared.",
- modelName,
- functionName);
- AiImageFunctionDef imageFunction =
findImageAiFunction(functionName);
- if (isTextAiFunction(functionName)) {
- Preconditions.checkArgument(
- modelClient instanceof SupportsTextGeneration,
- "Model '%s' referenced by %s does not support text
generation.",
- modelName,
- functionName);
- } else if (imageFunction != null) {
- validateImageAiModelCapability(
- imageFunction, modelClient, modelName,
functionName);
- } else {
- Preconditions.checkArgument(
- modelClient instanceof SupportsEmbedding,
- "Model '%s' referenced by %s does not support
embedding.",
- modelName,
- functionName);
- }
- }
- for (SqlNode operand : call.getOperandList()) {
- if (operand != null) {
- validateAiModelCapabilities(operand, modelClients,
userDefinedFunctionNames);
- }
- }
- } else if (node instanceof SqlNodeList) {
- for (SqlNode child : (SqlNodeList) node) {
- validateAiModelCapabilities(child, modelClients,
userDefinedFunctionNames);
- }
- }
- }
-
- private static void validateImageAiModelCapability(
- AiImageFunctionDef function,
- AiModelClient modelClient,
- String modelName,
- String functionName) {
- switch (function.getCapability()) {
- case IMAGE_TEXT_GENERATION:
- Preconditions.checkArgument(
- modelClient instanceof SupportsImageTextGeneration,
- "Model '%s' referenced by %s does not support image
text generation.",
- modelName,
- functionName);
- break;
- case IMAGE_EMBEDDING:
- Preconditions.checkArgument(
- modelClient instanceof SupportsImageEmbedding,
- "Model '%s' referenced by %s does not support image
embedding.",
- modelName,
- functionName);
- break;
- default:
- throw new IllegalArgumentException(
- "Unsupported capability for image AI function " +
functionName);
- }
- }
-
- private static String resolveAiModelName(SqlCall call) {
- SqlNode modelArgument = call.operand(0);
- Preconditions.checkArgument(
- modelArgument instanceof SqlCharStringLiteral,
- "The model argument of %s must be a string constant, but was
%s.",
- call.getOperator().getName(),
- modelArgument);
- return ((SqlCharStringLiteral)
modelArgument).getNlsString().getValue();
- }
-
- private static boolean isUserDefinedFunction(
- String functionName, Set<String> userDefinedFunctionNames) {
- return
userDefinedFunctionNames.stream().anyMatch(functionName::equalsIgnoreCase);
- }
-
- private static boolean isAiFunction(String functionName) {
- return isTextAiFunction(functionName)
- || isEmbeddingAiFunction(functionName)
- || findImageAiFunction(functionName) != null;
- }
-
- private static boolean isTextAiFunction(String functionName) {
- for (AiTextFunctionDef function : AiTextFunctionDef.values()) {
- if (function.getFunctionName().equalsIgnoreCase(functionName)) {
- return true;
- }
- }
- return false;
- }
-
- private static boolean isEmbeddingAiFunction(String functionName) {
- for (AiEmbeddingFunctionDef function :
AiEmbeddingFunctionDef.values()) {
- if (function.getFunctionName().equalsIgnoreCase(functionName)) {
- return true;
- }
- }
- return false;
- }
-
- @Nullable
- private static AiImageFunctionDef findImageAiFunction(String functionName)
{
- for (AiImageFunctionDef function : AiImageFunctionDef.values()) {
- if (function.getFunctionName().equalsIgnoreCase(functionName)) {
- return function;
- }
- }
- return null;
- }
-
public static boolean hasAsterisk(@Nullable String projection) {
if (isNullOrWhitespaceOnly(projection)) {
// Providing an empty projection expression is equivalent to
writing `*` explicitly.
diff --git
a/flink-cdc-runtime/src/test/java/org/apache/flink/cdc/runtime/functions/impl/AiFunctionsTest.java
b/flink-cdc-runtime/src/test/java/org/apache/flink/cdc/runtime/functions/impl/AiFunctionsTest.java
index 8eee49ae56..eac3612283 100644
---
a/flink-cdc-runtime/src/test/java/org/apache/flink/cdc/runtime/functions/impl/AiFunctionsTest.java
+++
b/flink-cdc-runtime/src/test/java/org/apache/flink/cdc/runtime/functions/impl/AiFunctionsTest.java
@@ -26,7 +26,9 @@ import
org.apache.flink.cdc.common.model.abilities.SupportsTextGeneration;
import org.junit.jupiter.api.Test;
import java.util.ArrayList;
+import java.util.Collections;
import java.util.List;
+import java.util.Map;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.assertThatThrownBy;
@@ -89,19 +91,23 @@ class AiFunctionsTest {
@Test
void testTextAiFunctionsUseEnglishPromptsAndParseJsonResponses() {
TestModelClient model = new TestModelClient();
+ Map<String, AiModelClient> modelClients = modelClients("textModel",
model);
- assertThat(AiFunctions.aiComplete(model, "input", "Return three
letters"))
+ assertThat(
+ AiFunctions.aiComplete(
+ "textModel", "input", "Return three letters",
modelClients))
.hasToString("{\"result\":\"ABC\"}");
- assertThat(AiFunctions.aiClassify(model, "input", "positive,negative"))
+ assertThat(AiFunctions.aiClassify("textModel", "input",
"positive,negative", modelClients))
.hasToString("{\"result\":\"ABC\"}");
- assertThat(AiFunctions.aiTranslate(model, "input", "auto", "en"))
+ assertThat(AiFunctions.aiTranslate("textModel", "input", "auto", "en",
modelClients))
.hasToString("{\"result\":\"ABC\"}");
- assertThat(AiFunctions.aiSummarize(model, "input", 100))
+ assertThat(AiFunctions.aiSummarize("textModel", "input", 100,
modelClients))
.hasToString("{\"result\":\"ABC\"}");
- assertThat(AiFunctions.aiSentiment(model,
"input")).hasToString("{\"result\":\"ABC\"}");
- assertThat(AiFunctions.aiExtract(model, "input", "name:string"))
+ assertThat(AiFunctions.aiSentiment("textModel", "input", modelClients))
.hasToString("{\"result\":\"ABC\"}");
- assertThat(AiFunctions.aiMask(model, "input", "email,phone"))
+ assertThat(AiFunctions.aiExtract("textModel", "input", "name:string",
modelClients))
+ .hasToString("{\"result\":\"ABC\"}");
+ assertThat(AiFunctions.aiMask("textModel", "input", "email,phone",
modelClients))
.hasToString("{\"result\":\"ABC\"}");
assertThat(model.prompts).hasSize(7);
@@ -128,19 +134,25 @@ class AiFunctionsTest {
@Test
void testEmbeddingFunction() {
TestModelClient model = new TestModelClient();
+ Map<String, AiModelClient> modelClients =
modelClients("embeddingModel", model);
- assertThat(AiFunctions.aiEmbed(model, "input")).containsExactly(0.1f,
0.2f, 0.3f);
+ assertThat(AiFunctions.aiEmbed("embeddingModel", "input",
modelClients))
+ .containsExactly(0.1f, 0.2f, 0.3f);
assertThat(model.embedCalls).isOne();
}
@Test
void testImageAiFunctions() {
TestModelClient model = new TestModelClient();
+ Map<String, AiModelClient> modelClients = modelClients("imageModel",
model);
byte[] image = new byte[] {1, 2, 3, 4};
- assertThat(AiFunctions.aiImageComplete(model, image, "Describe the
image"))
+ assertThat(
+ AiFunctions.aiImageComplete(
+ "imageModel", image, "Describe the image",
modelClients))
.isEqualTo("image has 4 bytes, prompt: Describe the image");
- assertThat(AiFunctions.aiImageEmbed(model,
image)).containsExactly(0.9f, 0.8f, 0.7f);
+ assertThat(AiFunctions.aiImageEmbed("imageModel", image, modelClients))
+ .containsExactly(0.9f, 0.8f, 0.7f);
assertThat(model.imageTextCalls).isOne();
assertThat(model.imageEmbedCalls).isOne();
}
@@ -148,26 +160,64 @@ class AiFunctionsTest {
@Test
void testUnsupportedCapabilities() {
UnsupportedModelClient model = new UnsupportedModelClient();
+ Map<String, AiModelClient> modelClients = modelClients("unsupported",
model);
- assertThatThrownBy(() -> AiFunctions.aiComplete(model, "input",
"prompt"))
+ assertThatThrownBy(
+ () ->
+ AiFunctions.aiComplete(
+ "unsupported", "input", "prompt",
modelClients))
.isInstanceOf(UnsupportedOperationException.class)
- .hasMessageContaining("does not support text generation");
- assertThatThrownBy(() -> AiFunctions.aiEmbed(model, "input"))
+ .hasMessageContaining("Model 'unsupported'")
+ .hasMessageContaining("AI_COMPLETE")
+ .hasMessageContaining("does not implement
SupportsTextGeneration interface");
+ assertThatThrownBy(() -> AiFunctions.aiEmbed("unsupported", "input",
modelClients))
.isInstanceOf(UnsupportedOperationException.class)
- .hasMessageContaining("does not support embedding");
- assertThatThrownBy(() -> AiFunctions.aiImageComplete(model, new byte[]
{1, 2}, "describe"))
+ .hasMessageContaining("Model 'unsupported'")
+ .hasMessageContaining("AI_EMBED")
+ .hasMessageContaining("does not implement SupportsEmbedding
interface");
+ assertThatThrownBy(
+ () ->
+ AiFunctions.aiImageComplete(
+ "unsupported", new byte[] {1, 2},
"describe", modelClients))
.isInstanceOf(UnsupportedOperationException.class)
- .hasMessageContaining("does not support image text
generation");
- assertThatThrownBy(() -> AiFunctions.aiImageEmbed(model, new byte[]
{1, 2}))
+ .hasMessageContaining("Model 'unsupported'")
+ .hasMessageContaining("AI_IMAGE_COMPLETE")
+ .hasMessageContaining("does not implement
SupportsImageTextGeneration interface");
+ assertThatThrownBy(
+ () ->
+ AiFunctions.aiImageEmbed(
+ "unsupported", new byte[] {1, 2},
modelClients))
.isInstanceOf(UnsupportedOperationException.class)
- .hasMessageContaining("does not support image embedding");
+ .hasMessageContaining("Model 'unsupported'")
+ .hasMessageContaining("AI_IMAGE_EMBED")
+ .hasMessageContaining("does not implement
SupportsImageEmbedding interface");
+ }
+
+ @Test
+ void testInvalidModelName() {
+ Map<String, AiModelClient> modelClients = Collections.emptyMap();
+
+ assertThatThrownBy(
+ () ->
+ AiFunctions.aiComplete(
+ "missingModel", "input", "prompt",
modelClients))
+ .isInstanceOf(IllegalArgumentException.class)
+ .hasMessage(
+ "Model 'missingModel' referenced by AI_COMPLETE has
not been declared.");
+ assertThatThrownBy(() -> AiFunctions.aiEmbed(null, "input",
modelClients))
+ .isInstanceOf(IllegalArgumentException.class)
+ .hasMessage("Model name referenced by AI_EMBED must not be
null.");
}
@Test
void testInvalidJsonResponse() {
TestModelClient model = new TestModelClient("not-json");
+ Map<String, AiModelClient> modelClients = modelClients("model", model);
- assertThatThrownBy(() -> AiFunctions.aiClassify(model, "input",
"positive,negative"))
+ assertThatThrownBy(
+ () ->
+ AiFunctions.aiClassify(
+ "model", "input", "positive,negative",
modelClients))
.isInstanceOf(RuntimeException.class)
.hasMessage("AI function AI_CLASSIFY returned invalid JSON:
not-json");
}
@@ -176,8 +226,12 @@ class AiFunctionsTest {
void testInvalidJsonResponseIsTruncated() {
String longInvalidJson = "x".repeat(600);
TestModelClient model = new TestModelClient(longInvalidJson);
+ Map<String, AiModelClient> modelClients = modelClients("model", model);
- assertThatThrownBy(() -> AiFunctions.aiClassify(model, "input",
"positive,negative"))
+ assertThatThrownBy(
+ () ->
+ AiFunctions.aiClassify(
+ "model", "input", "positive,negative",
modelClients))
.isInstanceOf(RuntimeException.class)
.hasMessage(
"AI function AI_CLASSIFY returned invalid JSON: "
@@ -188,11 +242,15 @@ class AiFunctionsTest {
@Test
void testNullInputSkipsModelInvocation() {
TestModelClient model = new TestModelClient();
-
- assertThat(AiFunctions.aiClassify(model, null,
"positive,negative")).isNull();
- assertThat(AiFunctions.aiEmbed(model, null)).isNull();
- assertThat(AiFunctions.aiImageComplete(model, null,
"describe")).isNull();
- assertThat(AiFunctions.aiImageEmbed(model, null)).isNull();
+ Map<String, AiModelClient> modelClients = modelClients("model", model);
+ Map<String, AiModelClient> emptyModelClients = Collections.emptyMap();
+
+ assertThat(AiFunctions.aiClassify("model", null, "positive,negative",
modelClients))
+ .isNull();
+ assertThat(AiFunctions.aiEmbed("model", null, modelClients)).isNull();
+ assertThat(AiFunctions.aiImageComplete("model", null, "describe",
modelClients)).isNull();
+ assertThat(AiFunctions.aiImageEmbed("model", null,
modelClients)).isNull();
+ assertThat(AiFunctions.aiComplete("missing", null, "prompt",
emptyModelClients)).isNull();
assertThat(model.prompts).isEmpty();
assertThat(model.embedCalls).isZero();
assertThat(model.imageTextCalls).isZero();
@@ -202,8 +260,13 @@ class AiFunctionsTest {
@Test
void testNullModelResponseReturnsNull() {
TestModelClient model = new TestModelClient(null);
+ Map<String, AiModelClient> modelClients = modelClients("model", model);
- assertThat(AiFunctions.aiSummarize(model, "input", 100)).isNull();
+ assertThat(AiFunctions.aiSummarize("model", "input", 100,
modelClients)).isNull();
assertThat(model.prompts).hasSize(1);
}
+
+ private static Map<String, AiModelClient> modelClients(String modelName,
AiModelClient model) {
+ return Map.of(modelName, model);
+ }
}
diff --git
a/flink-cdc-runtime/src/test/java/org/apache/flink/cdc/runtime/parser/AiFunctionParserTest.java
b/flink-cdc-runtime/src/test/java/org/apache/flink/cdc/runtime/parser/AiFunctionParserTest.java
index b6dc28f1ca..6f60d0883e 100644
---
a/flink-cdc-runtime/src/test/java/org/apache/flink/cdc/runtime/parser/AiFunctionParserTest.java
+++
b/flink-cdc-runtime/src/test/java/org/apache/flink/cdc/runtime/parser/AiFunctionParserTest.java
@@ -17,11 +17,6 @@
package org.apache.flink.cdc.runtime.parser;
-import org.apache.flink.cdc.common.model.AiModelClient;
-import org.apache.flink.cdc.common.model.abilities.SupportsEmbedding;
-import org.apache.flink.cdc.common.model.abilities.SupportsImageEmbedding;
-import org.apache.flink.cdc.common.model.abilities.SupportsImageTextGeneration;
-import org.apache.flink.cdc.common.model.abilities.SupportsTextGeneration;
import org.apache.flink.cdc.common.schema.Column;
import org.apache.flink.cdc.common.source.SupportedMetadataColumn;
import org.apache.flink.cdc.common.types.DataTypes;
@@ -32,11 +27,8 @@ import org.junit.jupiter.api.Test;
import java.util.Collections;
import java.util.List;
-import java.util.Map;
-import java.util.Set;
import static org.assertj.core.api.Assertions.assertThat;
-import static org.assertj.core.api.Assertions.assertThatCode;
import static org.assertj.core.api.Assertions.assertThatThrownBy;
/** Parser and Janino tests for the generic AI functions. */
@@ -48,44 +40,6 @@ class AiFunctionParserTest {
Column.physicalColumn("content", DataTypes.STRING()),
Column.physicalColumn("image", DataTypes.BYTES()));
- private static class TextModelClient implements AiModelClient,
SupportsTextGeneration {
- private static final long serialVersionUID = 1L;
-
- @Override
- public String generate(String systemPrompt, String userInput) {
- return "{}";
- }
- }
-
- private static class EmbeddingModelClient implements AiModelClient,
SupportsEmbedding {
- private static final long serialVersionUID = 1L;
-
- @Override
- public float[] embed(String text) {
- return new float[0];
- }
- }
-
- private static class ImageTextModelClient
- implements AiModelClient, SupportsImageTextGeneration {
- private static final long serialVersionUID = 1L;
-
- @Override
- public String generateTextFromImage(byte[] image, String prompt) {
- return "description";
- }
- }
-
- private static class ImageEmbeddingModelClient
- implements AiModelClient, SupportsImageEmbedding {
- private static final long serialVersionUID = 1L;
-
- @Override
- public float[] embedImage(byte[] image) {
- return new float[0];
- }
- }
-
@Test
void testTranslateAiFunctions() {
List<ProjectionColumn> columns =
@@ -96,7 +50,8 @@ class AiFunctionParserTest {
assertThat(columns)
.extracting(ProjectionColumn::getScriptExpression)
.containsExactly(
- "aiComplete(completer, $0, \"You are helpful\")",
"aiEmbed(embedder, $0)");
+ "aiComplete(\"completer\", $0, \"You are helpful\",
__ai_model_clients__)",
+ "aiEmbed(\"embedder\", $0, __ai_model_clients__)");
assertThat(columns)
.extracting(ProjectionColumn::getDataType)
.containsExactly(DataTypes.VARIANT(),
DataTypes.ARRAY(DataTypes.FLOAT()));
@@ -116,12 +71,18 @@ class AiFunctionParserTest {
assertThat(columns)
.extracting(ProjectionColumn::getScriptExpression)
.containsExactly(
- "aiClassify(model, $0, \"positive,negative\")",
- "aiTranslate(model, $0, \"auto\", \"en\")",
- "aiSummarize(model, $0, 100)",
- "aiSentiment(model, $0)",
- "aiExtract(model, $0, \"name:string\")",
- "aiMask(model, $0, \"email,phone\")");
+ "aiClassify(\"model\", $0, \"positive,negative\",
__ai_model_clients__)",
+ "aiTranslate(\n"
+ + " \"model\",\n"
+ + " $0,\n"
+ + " \"auto\",\n"
+ + " \"en\",\n"
+ + " __ai_model_clients__\n"
+ + ")",
+ "aiSummarize(\"model\", $0, 100,
__ai_model_clients__)",
+ "aiSentiment(\"model\", $0, __ai_model_clients__)",
+ "aiExtract(\"model\", $0, \"name:string\",
__ai_model_clients__)",
+ "aiMask(\"model\", $0, \"email,phone\",
__ai_model_clients__)");
assertThat(columns)
.extracting(ProjectionColumn::getDataType)
.containsOnly(DataTypes.VARIANT());
@@ -137,33 +98,30 @@ class AiFunctionParserTest {
assertThat(columns)
.extracting(ProjectionColumn::getScriptExpression)
.containsExactly(
- "aiImageComplete(vision, $0, \"Describe the image\")",
- "aiImageEmbed(imageEmbedder, $0)");
+ "aiImageComplete(\"vision\", $0, \"Describe the
image\", __ai_model_clients__)",
+ "aiImageEmbed(\"imageEmbedder\", $0,
__ai_model_clients__)");
assertThat(columns)
.extracting(ProjectionColumn::getDataType)
.containsExactly(DataTypes.STRING(),
DataTypes.ARRAY(DataTypes.FLOAT()));
}
@Test
- void testSameNamedUdfTakesPrecedenceOverAiFunction() {
- Set<String> udfNames = Set.of("ai_sentiment");
- assertThatCode(
- () ->
- TransformParser.validateAiModelReferences(
- "AI_SENTIMENT(id) AS sentiment",
- null,
- Collections.emptySet(),
- udfNames))
- .doesNotThrowAnyException();
- assertThatCode(
- () ->
- TransformParser.validateAiModelCapabilities(
- "AI_SENTIMENT(id) AS sentiment",
- null,
- Collections.emptyMap(),
- udfNames))
- .doesNotThrowAnyException();
+ void testDynamicModelSelection() {
+ List<ProjectionColumn> columns =
+ translate(
+ "AI_COMPLETE(IF(id = 1, 'powerful', 'cheap'), content,
'prompt') AS completed");
+ assertThat(columns)
+ .extracting(ProjectionColumn::getScriptExpression)
+ .containsExactly(
+ "aiComplete(isTrue(valueEquals($0, 1)) ? \"powerful\"
: \"cheap\", $1, \"prompt\", __ai_model_clients__)");
+ assertThat(columns)
+ .extracting(ProjectionColumn::getDataType)
+ .containsExactly(DataTypes.VARIANT());
+ }
+
+ @Test
+ void testSameNamedUdfTakesPrecedenceOverAiFunction() {
List<ProjectionColumn> columns =
TransformParser.generateProjectionColumns(
"AI_SENTIMENT(id) AS sentiment",
@@ -192,18 +150,6 @@ class AiFunctionParserTest {
private static void assertSameNamedUdfTakesPrecedenceOverImageAiFunction(
String functionName, String udfName) {
String projection = functionName + "(id) AS udf_output";
- Set<String> udfNames = Set.of(udfName);
-
- assertThatCode(
- () ->
- TransformParser.validateAiModelReferences(
- projection, null,
Collections.emptySet(), udfNames))
- .doesNotThrowAnyException();
- assertThatCode(
- () ->
- TransformParser.validateAiModelCapabilities(
- projection, null,
Collections.emptyMap(), udfNames))
- .doesNotThrowAnyException();
List<ProjectionColumn> columns =
TransformParser.generateProjectionColumns(
@@ -223,110 +169,6 @@ class AiFunctionParserTest {
.containsExactly(DataTypes.STRING());
}
- @Test
- void testModelArgumentMustBeStringConstant() {
- assertThatThrownBy(
- () ->
- TransformParser.validateAiModelReferences(
- "AI_COMPLETE(content, content,
'prompt') AS completed",
- null,
- Set.of("content")))
- .isInstanceOf(IllegalArgumentException.class)
- .hasMessageContaining("must be a string constant");
- }
-
- @Test
- void testReferencedModelMustBeDeclared() {
- assertThatThrownBy(
- () ->
- TransformParser.validateAiModelReferences(
- "AI_EMBED('missing', content) AS
embedding",
- null,
- Set.of("declared")))
- .isInstanceOf(IllegalArgumentException.class)
- .hasMessageContaining("Model 'missing'")
- .hasMessageContaining("has not been declared");
-
- assertThatCode(
- () ->
- TransformParser.validateAiModelReferences(
- "AI_EMBED('declared', content) AS
embedding",
- null,
- Set.of("declared")))
- .doesNotThrowAnyException();
-
- assertThatThrownBy(
- () ->
- TransformParser.validateAiModelReferences(
- "AI_CLASSIFY('missing', content,
'a,b') AS classified",
- null,
- Set.of("declared")))
- .isInstanceOf(IllegalArgumentException.class)
- .hasMessageContaining("Model 'missing'")
- .hasMessageContaining("AI_CLASSIFY");
- }
-
- @Test
- void testModelCapabilitiesMustMatchAiFunctions() {
- Map<String, AiModelClient> models =
- Map.of(
- "textModel", new TextModelClient(),
- "embeddingModel", new EmbeddingModelClient(),
- "imageTextModel", new ImageTextModelClient(),
- "imageEmbeddingModel", new
ImageEmbeddingModelClient());
-
- assertThatCode(
- () ->
- TransformParser.validateAiModelCapabilities(
- "AI_CLASSIFY('textModel', content,
'a,b') AS classified, "
- + "AI_EMBED('embeddingModel',
content) AS embedding, "
- +
"AI_IMAGE_COMPLETE('imageTextModel', image, 'describe') AS description, "
- +
"AI_IMAGE_EMBED('imageEmbeddingModel', image) AS image_embedding",
- null,
- models))
- .doesNotThrowAnyException();
- assertThatThrownBy(
- () ->
- TransformParser.validateAiModelCapabilities(
- "AI_SENTIMENT('embeddingModel',
content) AS sentiment",
- null,
- models))
- .isInstanceOf(IllegalArgumentException.class)
- .hasMessageContaining("Model 'embeddingModel'")
- .hasMessageContaining("AI_SENTIMENT")
- .hasMessageContaining("does not support text generation");
- assertThatThrownBy(
- () ->
- TransformParser.validateAiModelCapabilities(
- "AI_EMBED('textModel', content) AS
embedding",
- null,
- models))
- .isInstanceOf(IllegalArgumentException.class)
- .hasMessageContaining("Model 'textModel'")
- .hasMessageContaining("AI_EMBED")
- .hasMessageContaining("does not support embedding");
- assertThatThrownBy(
- () ->
- TransformParser.validateAiModelCapabilities(
- "AI_IMAGE_COMPLETE('textModel', image,
'describe') AS description",
- null,
- models))
- .isInstanceOf(IllegalArgumentException.class)
- .hasMessageContaining("Model 'textModel'")
- .hasMessageContaining("AI_IMAGE_COMPLETE")
- .hasMessageContaining("does not support image text
generation");
- assertThatThrownBy(
- () ->
- TransformParser.validateAiModelCapabilities(
- "AI_IMAGE_EMBED('embeddingModel',
image) AS embedding",
- null,
- models))
- .isInstanceOf(IllegalArgumentException.class)
- .hasMessageContaining("Model 'embeddingModel'")
- .hasMessageContaining("AI_IMAGE_EMBED")
- .hasMessageContaining("does not support image embedding");
- }
-
@Test
void testFunctionArityValidation() {
assertThatThrownBy(() -> translate("AI_EMBED('model') AS embedding"))
@@ -352,11 +194,6 @@ class AiFunctionParserTest {
"Invalid number of arguments to function
'AI_IMAGE_COMPLETE'");
assertThatThrownBy(() -> translate("AI_IMAGE_EMBED('model') AS
embedding"))
.hasMessageContaining("Invalid number of arguments to function
'AI_IMAGE_EMBED'");
- assertThatCode(
- () ->
- TransformParser.validateAiModelReferences(
- "AI_COMPLETE() AS completed", null,
Collections.emptySet()))
- .doesNotThrowAnyException();
}
private List<ProjectionColumn> translate(String expression) {