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) {

Reply via email to