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 f535d56eff [FLINK-40416][runtime] Add built-in text AI functions 
(#4512)
f535d56eff is described below

commit f535d56eff62b9d62403c1d015db075ca17f2f18
Author: haruki <[email protected]>
AuthorDate: Tue Sep 1 15:43:23 2026 +0800

    [FLINK-40416][runtime] Add built-in text AI functions (#4512)
---
 docs/content.zh/docs/core-concept/ai-model.md      |  28 +++-
 docs/content/docs/core-concept/ai-model.md         |  28 +++-
 .../flink/translator/TransformTranslator.java      |  43 +++++-
 .../flink/FlinkPipelineAiFunctionITCase.java       |  49 ++++++-
 .../cdc/pipeline/tests/AiFunctionE2eITCase.java    |  30 +++--
 .../flink/cdc/models/dummy/DummyModelClient.java   |  18 +++
 .../flink/cdc/runtime/ai/AiTextFunctionDef.java    |  65 ++++++++-
 .../cdc/runtime/functions/impl/AiFunctions.java    |  62 ++++++++-
 .../flink/cdc/runtime/parser/TransformParser.java  | 145 ++++++++++++++++++---
 .../runtime/functions/impl/AiFunctionsTest.java    | 104 +++++++++++++--
 .../cdc/runtime/parser/AiFunctionParserTest.java   | 145 +++++++++++++++++++++
 11 files changed, 658 insertions(+), 59 deletions(-)

diff --git a/docs/content.zh/docs/core-concept/ai-model.md 
b/docs/content.zh/docs/core-concept/ai-model.md
index 9b80293265..0a36d17fa5 100644
--- a/docs/content.zh/docs/core-concept/ai-model.md
+++ b/docs/content.zh/docs/core-concept/ai-model.md
@@ -24,11 +24,30 @@ under the License.
 
 # AI 模型
 
-AI 模型可用于 transform 表达式中的文本补全和 embedding。
+AI 模型可用于 transform 表达式中的文本生成、文本分析和 embedding。
+
+## AI Functions
+
+模型名称必须是字符串常量,并引用 `pipeline.model` 中声明的模型。文本函数要求模型客户端支持文本生成,`AI_EMBED` 
要求模型客户端支持 embedding;Pipeline 会在执行前校验引用模型的 capability 是否匹配。
+
+所有文本函数都会将模型返回的 JSON 解析为 `VARIANT`。
+
+| 函数 | 说明 | JSON 字段 |
+|------|------|-----------|
+| `AI_COMPLETE(model, input, system_prompt)` | 使用调用方提供的 system prompt 补全文本。 | 
`result` |
+| `AI_CLASSIFY(model, input, labels)` | 将输入分类到给定标签之一。 | 
`category`、`confidence` |
+| `AI_TRANSLATE(model, input, source_lang, target_lang)` | 翻译输入;`source_lang` 
可传 `auto` 自动识别。 | `translated_text`、`detected_language` |
+| `AI_SUMMARIZE(model, input, max_length)` | 在指定字符数内生成摘要。 | `summary` |
+| `AI_SENTIMENT(model, input)` | 分析文本情感。 | `score`、`label`、`confidence` |
+| `AI_EXTRACT(model, input, schema)` | 按 schema 字符串描述提取字段。 | `extracted_json` |
+| `AI_MASK(model, input, entities)` | 对指定实体类型进行脱敏。 | 
`masked_text`、`detected_entities` |
+| `AI_EMBED(model, input)` | 生成 embedding 向量。 | 不返回 JSON,而是返回 `ARRAY<FLOAT>`。 |
+
+六个专用文本函数使用内置英文 prompt 模板,但输入文本可以是任意语言。输入为 `NULL` 时直接返回 `NULL`,且不会调用模型;模型返回 
`NULL` 时也返回 `NULL`。非空文本响应必须是语法合法的 JSON,否则当前记录处理失败,错误信息会标明具体 AI 函数。运行时只校验 JSON 
语法,不校验响应字段是否存在或字段类型是否匹配。
 
 ## OpenAI-compatible 模型客户端
 
-AI 模型客户端可供 transform 中的 `AI_COMPLETE` 和 `AI_EMBED` 函数引用。使用时,需要通过 `--jar` 将模型实现 
JAR(例如 `flink-cdc-pipeline-model-openai-compatible`)添加到 Pipeline 命令中。
+AI 模型客户端可供上述 AI Functions 引用。使用时,需要通过 `--jar` 将模型实现 JAR(例如 
`flink-cdc-pipeline-model-openai-compatible`)添加到 Pipeline 命令中。
 
 OpenAI-compatible 客户端支持调用实现 OpenAI Chat Completions 和 Embeddings REST API 的服务。
 
@@ -40,6 +59,7 @@ transform:
     projection: >-
       *,
       AI_COMPLETE('completion_model', content, '总结输入内容') AS summary,
+      AI_SENTIMENT('completion_model', content) AS sentiment,
       AI_EMBED('embedding_model', content) AS embedding
 
 pipeline:
@@ -71,11 +91,11 @@ pipeline:
 | `model` | 是 | 发送给服务端的模型名称;`model-name` 作为废弃别名仍可使用。 |
 | `endpoint` | 是 | OpenAI-compatible 服务的 Base URL。 |
 | `api-key` | 是 | 请求认证使用的 Bearer Token。 |
-| `system-prompt` | 否 | 添加在 `AI_COMPLETE` 生成的 system prompt 之前。 |
+| `system-prompt` | 否 | 添加在所有文本 AI Function 的 prompt 之前。 |
 | `user-prompt` | 否 | 在输入之后追加一条 user message。 |
 | `temperature`、`top-p`、`stop`、`max-tokens` | 否 | 常用文本生成参数。 |
 | `presence-penalty`、`frequency-penalty`、`n`、`seed` | 否 | 其他文本生成参数。 |
-| `response-format` | 否 | 支持 `json_object`;AI completion 的结果必须是合法 JSON。 |
+| `response-format` | 否 | 支持 `json_object`;文本 AI Function 的结果必须是合法 JSON。 |
 | `content-type` | 否 | `text`(默认)或 `image_url`。 |
 | `dimension` | 否 | 请求的 embedding 维度。 |
 | `extra-header`、`extra-body` | 否 | JSON 对象格式的厂商自定义请求头或请求体字段。 |
diff --git a/docs/content/docs/core-concept/ai-model.md 
b/docs/content/docs/core-concept/ai-model.md
index da3ef0f008..8aa7666838 100644
--- a/docs/content/docs/core-concept/ai-model.md
+++ b/docs/content/docs/core-concept/ai-model.md
@@ -24,11 +24,30 @@ under the License.
 
 # AI Model
 
-AI models can be used in transform expressions for text completion and 
embedding.
+AI models can be used in transform expressions for text generation, text 
analysis, and embedding.
+
+## 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 `AI_EMBED` requires embedding support. The pipeline validates 
the referenced model capability before execution.
+
+All text functions return `VARIANT` values parsed from the model's JSON 
response.
+
+| Function | Description | JSON fields |
+|----------|-------------|-------------|
+| `AI_COMPLETE(model, input, system_prompt)` | Completes the input using a 
caller-provided system prompt. | `result` |
+| `AI_CLASSIFY(model, input, labels)` | Classifies the input into one of the 
provided labels. | `category`, `confidence` |
+| `AI_TRANSLATE(model, input, source_lang, target_lang)` | Translates the 
input. Use `auto` to detect the source language. | `translated_text`, 
`detected_language` |
+| `AI_SUMMARIZE(model, input, max_length)` | Summarizes the input within the 
requested character limit. | `summary` |
+| `AI_SENTIMENT(model, input)` | Analyzes sentiment. | `score`, `label`, 
`confidence` |
+| `AI_EXTRACT(model, input, schema)` | Extracts fields described by the schema 
string. | `extracted_json` |
+| `AI_MASK(model, input, entities)` | Masks the requested entity types. | 
`masked_text`, `detected_entities` |
+| `AI_EMBED(model, input)` | Creates an embedding vector. | Returns 
`ARRAY<FLOAT>` instead of JSON. |
+
+The specialized text functions use built-in English prompt templates, but 
their input may be in any language. If the input is `NULL`, the function 
returns `NULL` without invoking the model. A `NULL` model response also 
produces `NULL`. A non-null text response must be syntactically valid JSON; 
otherwise, record processing fails with an error that identifies the AI 
function. The runtime validates JSON syntax but does not validate the presence 
or types of individual response fields.
 
 ## OpenAI-compatible Model Client
 
-AI model clients can be referenced by the `AI_COMPLETE` and `AI_EMBED` 
transform functions. Add the model implementation JAR, such as 
`flink-cdc-pipeline-model-openai-compatible`, to the pipeline command with 
`--jar`.
+AI model clients can be referenced by the AI functions above. Add the model 
implementation JAR, such as `flink-cdc-pipeline-model-openai-compatible`, to 
the pipeline command with `--jar`.
 
 The OpenAI-compatible client supports chat completions and text embeddings 
against endpoints that implement the corresponding OpenAI REST APIs.
 
@@ -40,6 +59,7 @@ transform:
     projection: >-
       *,
       AI_COMPLETE('completion_model', content, 'Summarize the input') AS 
summary,
+      AI_SENTIMENT('completion_model', content) AS sentiment,
       AI_EMBED('embedding_model', content) AS embedding
 
 pipeline:
@@ -71,11 +91,11 @@ Do not store API keys in source control. Supply them 
through the secret-manageme
 | `model` | Yes | Model name sent to the endpoint. `model-name` is accepted as 
a deprecated alias. |
 | `endpoint` | Yes | Base URL of the OpenAI-compatible endpoint. |
 | `api-key` | Yes | Bearer token used to authenticate requests. |
-| `system-prompt` | No | Prompt prepended to the system prompt generated by 
`AI_COMPLETE`. |
+| `system-prompt` | No | Prompt prepended to every text AI function prompt. |
 | `user-prompt` | No | Additional user message appended after the input. |
 | `temperature`, `top-p`, `stop`, `max-tokens` | No | Common generation 
parameters. |
 | `presence-penalty`, `frequency-penalty`, `n`, `seed` | No | Additional 
generation parameters. |
-| `response-format` | No | `json_object` is supported. AI completion results 
must be valid JSON. |
+| `response-format` | No | `json_object` is supported. Text AI function 
results must be valid JSON. |
 | `content-type` | No | `text` (default) or `image_url`. |
 | `dimension` | No | Requested embedding dimension. |
 | `extra-header`, `extra-body` | No | Provider-specific headers or body fields 
encoded as JSON objects. |
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 20c8bc54ee..5dc9b268f7 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
@@ -41,6 +41,7 @@ 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;
@@ -66,7 +67,8 @@ public class TransformTranslator {
         if (transforms.isEmpty()) {
             return input;
         }
-        validateModelReferences(transforms, models);
+        validateModelReferences(
+                transforms, models, getUserDefinedFunctionNames(udfFunctions, 
models));
         return input.transform(
                 "Transform:Schema",
                 new EventTypeInfo(),
@@ -144,7 +146,10 @@ public class TransformTranslator {
                         .filter(ModelDef::isLegacy)
                         .map(this::modelToUDFTuple)
                         .collect(Collectors.toList()));
-        postTransformFunctionBuilder.addModelClients(loadModelClients(models, 
env));
+        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())
                 .uid(operatorUidGenerator.generateUid("post-transform"));
@@ -186,7 +191,10 @@ public class TransformTranslator {
         return clients;
     }
 
-    private void validateModelReferences(List<TransformDef> transforms, 
List<ModelDef> models) {
+    private void validateModelReferences(
+            List<TransformDef> transforms,
+            List<ModelDef> models,
+            Set<String> userDefinedFunctionNames) {
         Set<String> clientModelNames =
                 models.stream()
                         .filter(model -> !model.isLegacy())
@@ -194,10 +202,37 @@ public class TransformTranslator {
                         .collect(Collectors.toSet());
         for (TransformDef transform : transforms) {
             TransformParser.validateAiModelReferences(
-                    transform.getProjection(), transform.getFilter(), 
clientModelNames);
+                    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 f01777540e..d50f1c67d1 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
@@ -34,6 +34,7 @@ import org.apache.flink.cdc.composer.definition.PipelineDef;
 import org.apache.flink.cdc.composer.definition.SinkDef;
 import org.apache.flink.cdc.composer.definition.SourceDef;
 import org.apache.flink.cdc.composer.definition.TransformDef;
+import org.apache.flink.cdc.composer.definition.UdfDef;
 import org.apache.flink.cdc.connectors.values.ValuesDatabase;
 import org.apache.flink.cdc.connectors.values.factory.ValuesDataFactory;
 import org.apache.flink.cdc.connectors.values.sink.ValuesDataSinkOptions;
@@ -127,6 +128,42 @@ class FlinkPipelineAiFunctionITCase {
                         "Dummy model closed.");
     }
 
+    @Test
+    void testSpecializedTextAiFunctionsInProjection() throws Exception {
+        String[] output =
+                runAiFunctionTest(
+                        "id, "
+                                + "AI_CLASSIFY('testModel', content, 
'positive,negative') AS classified, "
+                                + "AI_TRANSLATE('testModel', content, 'auto', 
'en') AS translated, "
+                                + "AI_SUMMARIZE('testModel', content, 100) AS 
summarized, "
+                                + "AI_SENTIMENT('testModel', content) AS 
sentiment, "
+                                + "AI_EXTRACT('testModel', content, 
'name:string') AS extracted, "
+                                + "AI_MASK('testModel', content, 'name') AS 
masked",
+                        List.of(ModelDef.of("testModel", "dummy", 
Collections.emptyMap())));
+
+        assertThat(output)
+                .containsExactly(
+                        
"CreateTableEvent{tableId=default_namespace.default_schema.mytable1, 
schema=columns={`id` INT NOT NULL,`classified` VARIANT,`translated` 
VARIANT,`summarized` VARIANT,`sentiment` VARIANT,`extracted` VARIANT,`masked` 
VARIANT}, primaryKeys=id, options=()}",
+                        
"DataChangeEvent{tableId=default_namespace.default_schema.mytable1, before=[], 
after=[1, {\"category\":\"dummy\",\"confidence\":1}, 
{\"detected_language\":\"en\",\"translated_text\":\"dummy translation\"}, 
{\"summary\":\"dummy summary\"}, 
{\"confidence\":1,\"label\":\"neutral\",\"score\":0}, 
{\"extracted_json\":{\"name\":\"dummy\"}}, 
{\"detected_entities\":\"name\",\"masked_text\":\"d***y\"}], op=INSERT, 
meta=()}");
+    }
+
+    @Test
+    void testSameNamedUdfTakesPrecedenceOverAiFunction() throws Exception {
+        String[] output =
+                runAiFunctionTest(
+                        "id, AI_SENTIMENT(id) AS sentiment",
+                        List.of(ModelDef.of("unusedModel", "dummy", 
Collections.emptyMap())),
+                        List.of(
+                                new UdfDef(
+                                        "ai_sentiment",
+                                        
"org.apache.flink.cdc.udf.examples.java.AddOneFunctionClass")));
+
+        assertThat(output)
+                .containsExactly(
+                        
"CreateTableEvent{tableId=default_namespace.default_schema.mytable1, 
schema=columns={`id` INT NOT NULL,`sentiment` STRING}, primaryKeys=id, 
options=()}",
+                        
"DataChangeEvent{tableId=default_namespace.default_schema.mytable1, before=[], 
after=[1, 2], op=INSERT, meta=()}");
+    }
+
     @Test
     void testAiEmbedInProjection() throws Exception {
         String[] output =
@@ -140,18 +177,24 @@ class FlinkPipelineAiFunctionITCase {
     }
 
     private String[] runAiFunctionTest(String projection, List<ModelDef> 
models) throws Exception {
+        return runAiFunctionTest(projection, models, Collections.emptyList());
+    }
+
+    private String[] runAiFunctionTest(
+            String projection, List<ModelDef> models, List<UdfDef> 
udfFunctions) throws Exception {
         URL modelJar = createDummyModelJar().toUri().toURL();
         ClassLoader originalClassLoader = 
Thread.currentThread().getContextClassLoader();
         try (URLClassLoader modelClassLoader =
                 new DummyModelClassLoader(modelJar, originalClassLoader)) {
             Thread.currentThread().setContextClassLoader(modelClassLoader);
-            return runAiFunctionTest(projection, models, modelJar);
+            return runAiFunctionTest(projection, models, udfFunctions, 
modelJar);
         } finally {
             Thread.currentThread().setContextClassLoader(originalClassLoader);
         }
     }
 
-    private String[] runAiFunctionTest(String projection, List<ModelDef> 
models, URL modelJar)
+    private String[] runAiFunctionTest(
+            String projection, List<ModelDef> models, List<UdfDef> 
udfFunctions, URL modelJar)
             throws Exception {
         FlinkPipelineComposer composer = FlinkPipelineComposer.ofMiniCluster();
 
@@ -212,7 +255,7 @@ class FlinkPipelineAiFunctionITCase {
                         sinkDef,
                         Collections.emptyList(),
                         Collections.singletonList(transformDef),
-                        Collections.emptyList(),
+                        udfFunctions,
                         models,
                         pipelineConfig);
 
diff --git 
a/flink-cdc-e2e-tests/flink-cdc-pipeline-e2e-tests/src/test/java/org/apache/flink/cdc/pipeline/tests/AiFunctionE2eITCase.java
 
b/flink-cdc-e2e-tests/flink-cdc-pipeline-e2e-tests/src/test/java/org/apache/flink/cdc/pipeline/tests/AiFunctionE2eITCase.java
index 8bdb7cb9a9..8264898974 100644
--- 
a/flink-cdc-e2e-tests/flink-cdc-pipeline-e2e-tests/src/test/java/org/apache/flink/cdc/pipeline/tests/AiFunctionE2eITCase.java
+++ 
b/flink-cdc-e2e-tests/flink-cdc-pipeline-e2e-tests/src/test/java/org/apache/flink/cdc/pipeline/tests/AiFunctionE2eITCase.java
@@ -30,7 +30,13 @@ class AiFunctionE2eITCase extends PipelineTestEnvironment {
 
     private static final String TABLE_1 = 
"default_namespace.default_schema.table1";
     private static final String TABLE_2 = 
"default_namespace.default_schema.table2";
-    private static final String DUMMY_JSON = "{\"result\":\"dummy response\"}";
+    private static final String DUMMY_TEXT_RESULTS =
+            "{\"category\":\"dummy\",\"confidence\":1}, "
+                    + 
"{\"detected_language\":\"en\",\"translated_text\":\"dummy translation\"}, "
+                    + "{\"summary\":\"dummy summary\"}, "
+                    + "{\"confidence\":1,\"label\":\"neutral\",\"score\":0}, "
+                    + "{\"extracted_json\":{\"name\":\"dummy\"}}, "
+                    + 
"{\"detected_entities\":\"name\",\"masked_text\":\"d***y\"}";
     private static final String DUMMY_EMBEDDING = "[3.0, 1.0, 4.0, 1.0, 5.0, 
9.0, 2.0, 6.0]";
 
     @Test
@@ -47,7 +53,13 @@ class AiFunctionE2eITCase extends PipelineTestEnvironment {
                         + "  - source-table: "
                         + TABLE_1
                         + "\n"
-                        + "    projection: col1, AI_COMPLETE('myModel', col1, 
'Complete it') AS completed\n"
+                        + "    projection: col1, "
+                        + "AI_CLASSIFY('myModel', col1, 'a,b') AS classified, "
+                        + "AI_TRANSLATE('myModel', col1, 'auto', 'en') AS 
translated, "
+                        + "AI_SUMMARIZE('myModel', col1, 100) AS summarized, "
+                        + "AI_SENTIMENT('myModel', col1) AS sentiment, "
+                        + "AI_EXTRACT('myModel', col1, 'name:string') AS 
extracted, "
+                        + "AI_MASK('myModel', col1, 'name') AS masked\n"
                         + "  - source-table: "
                         + TABLE_2
                         + "\n"
@@ -70,33 +82,33 @@ class AiFunctionE2eITCase extends PipelineTestEnvironment {
         validateResult(
                 "CreateTableEvent{tableId="
                         + TABLE_1
-                        + ", schema=columns={`col1` STRING NOT 
NULL,`completed` VARIANT}, primaryKeys=col1, options=()}",
+                        + ", schema=columns={`col1` STRING NOT 
NULL,`classified` VARIANT,`translated` VARIANT,`summarized` VARIANT,`sentiment` 
VARIANT,`extracted` VARIANT,`masked` VARIANT}, primaryKeys=col1, options=()}",
                 "DataChangeEvent{tableId="
                         + TABLE_1
                         + ", before=[], after=[1, "
-                        + DUMMY_JSON
+                        + DUMMY_TEXT_RESULTS
                         + "], op=INSERT, meta=()}",
                 "DataChangeEvent{tableId="
                         + TABLE_1
                         + ", before=[], after=[2, "
-                        + DUMMY_JSON
+                        + DUMMY_TEXT_RESULTS
                         + "], op=INSERT, meta=()}",
                 "DataChangeEvent{tableId="
                         + TABLE_1
                         + ", before=[], after=[3, "
-                        + DUMMY_JSON
+                        + DUMMY_TEXT_RESULTS
                         + "], op=INSERT, meta=()}",
                 "DataChangeEvent{tableId="
                         + TABLE_1
                         + ", before=[1, "
-                        + DUMMY_JSON
+                        + DUMMY_TEXT_RESULTS
                         + "], after=[], op=DELETE, meta=()}",
                 "DataChangeEvent{tableId="
                         + TABLE_1
                         + ", before=[2, "
-                        + DUMMY_JSON
+                        + DUMMY_TEXT_RESULTS
                         + "], after=[2, "
-                        + DUMMY_JSON
+                        + DUMMY_TEXT_RESULTS
                         + "], op=UPDATE, meta=()}");
 
         validateResult(
diff --git 
a/flink-cdc-pipeline-model/flink-cdc-pipeline-model-dummy/src/main/java/org/apache/flink/cdc/models/dummy/DummyModelClient.java
 
b/flink-cdc-pipeline-model/flink-cdc-pipeline-model-dummy/src/main/java/org/apache/flink/cdc/models/dummy/DummyModelClient.java
index f4c2b7444f..d92c0a8afb 100644
--- 
a/flink-cdc-pipeline-model/flink-cdc-pipeline-model-dummy/src/main/java/org/apache/flink/cdc/models/dummy/DummyModelClient.java
+++ 
b/flink-cdc-pipeline-model/flink-cdc-pipeline-model-dummy/src/main/java/org/apache/flink/cdc/models/dummy/DummyModelClient.java
@@ -37,6 +37,24 @@ public class DummyModelClient implements AiModelClient, 
SupportsTextGeneration,
         if (debug) {
             System.out.printf("Received prompt: %s%nUser input: %s%n", 
systemPrompt, userInput);
         }
+        if (systemPrompt.contains("\"category\"")) {
+            return "{\"category\":\"dummy\",\"confidence\":1.0}";
+        }
+        if (systemPrompt.contains("\"translated_text\"")) {
+            return "{\"translated_text\":\"dummy 
translation\",\"detected_language\":\"en\"}";
+        }
+        if (systemPrompt.contains("\"summary\"")) {
+            return "{\"summary\":\"dummy summary\"}";
+        }
+        if (systemPrompt.contains("\"score\"")) {
+            return "{\"score\":0.0,\"label\":\"neutral\",\"confidence\":1.0}";
+        }
+        if (systemPrompt.contains("\"extracted_json\"")) {
+            return "{\"extracted_json\":{\"name\":\"dummy\"}}";
+        }
+        if (systemPrompt.contains("\"masked_text\"")) {
+            return 
"{\"masked_text\":\"d***y\",\"detected_entities\":\"name\"}";
+        }
         return "{\"result\":\"dummy response\"}";
     }
 
diff --git 
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/ai/AiTextFunctionDef.java
 
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/ai/AiTextFunctionDef.java
index 74687617af..cbb8fc86cd 100644
--- 
a/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/ai/AiTextFunctionDef.java
+++ 
b/flink-cdc-runtime/src/main/java/org/apache/flink/cdc/runtime/ai/AiTextFunctionDef.java
@@ -21,13 +21,74 @@ import org.apache.flink.cdc.common.types.DataType;
 import org.apache.flink.cdc.common.types.DataTypes;
 import org.apache.flink.cdc.common.types.RowType;
 
+import java.util.Locale;
+
 /** Built-in AI text generation function definitions. */
 public enum AiTextFunctionDef {
     AI_COMPLETE(
             "AI_COMPLETE",
             RowType.of(new DataType[] {DataTypes.STRING()}, new String[] 
{"systemPrompt"}),
             RowType.of(new DataType[] {DataTypes.STRING()}, new String[] 
{"result"}),
-            "%s\n");
+            "%s\n"),
+
+    AI_CLASSIFY(
+            "AI_CLASSIFY",
+            RowType.of(new DataType[] {DataTypes.STRING()}, new String[] 
{"labels"}),
+            RowType.of(
+                    new DataType[] {DataTypes.STRING(), DataTypes.DOUBLE()},
+                    new String[] {"category", "confidence"}),
+            "You are a text classifier. Classify the input into exactly one of 
these labels: %s.\n"
+                    + "Choose only a provided label. Use the dominant meaning 
when multiple labels apply, "
+                    + "and lower the confidence when no label is a good 
match.\n"),
+
+    AI_TRANSLATE(
+            "AI_TRANSLATE",
+            RowType.of(
+                    new DataType[] {DataTypes.STRING(), DataTypes.STRING()},
+                    new String[] {"sourceLang", "targetLang"}),
+            RowType.of(
+                    new DataType[] {DataTypes.STRING(), DataTypes.STRING()},
+                    new String[] {"translated_text", "detected_language"}),
+            "You are a translator. Translate the input from %s to %s while 
preserving its meaning, "
+                    + "formatting, and terminology. If the source language is 
auto, detect it and report "
+                    + "the detected language code.\n"),
+
+    AI_SUMMARIZE(
+            "AI_SUMMARIZE",
+            RowType.of(new DataType[] {DataTypes.INT()}, new String[] 
{"maxLength"}),
+            RowType.of(new DataType[] {DataTypes.STRING()}, new String[] 
{"summary"}),
+            "You are a text summarizer. Summarize the input in no more than %d 
characters. Preserve "
+                    + "the key facts and conclusions, remove redundancy, and 
avoid subjective commentary.\n"),
+
+    AI_SENTIMENT(
+            "AI_SENTIMENT",
+            RowType.of(new DataType[0], new String[0]),
+            RowType.of(
+                    new DataType[] {DataTypes.DOUBLE(), DataTypes.STRING(), 
DataTypes.DOUBLE()},
+                    new String[] {"score", "label", "confidence"}),
+            "You are a sentiment analyzer. Analyze the input in context. 
Return a score from -1.0 "
+                    + "(most negative) to 1.0 (most positive), a label of 
positive, negative, or neutral, "
+                    + "and a confidence from 0.0 to 1.0. Consider tone, 
negation, and sarcasm.\n"),
+
+    AI_EXTRACT(
+            "AI_EXTRACT",
+            RowType.of(new DataType[] {DataTypes.STRING()}, new String[] 
{"schema"}),
+            RowType.of(new DataType[] {DataTypes.STRING()}, new String[] 
{"extracted_json"}),
+            "You are an information extraction system. Extract information 
from the input according "
+                    + "to this schema: %s. Preserve the requested field names 
and types, use null for "
+                    + "missing values, and place the extracted JSON object in 
extracted_json. Supported "
+                    + "types include string, number, integer, boolean, array, 
object, date, datetime, "
+                    + "email, and phone.\n"),
+
+    AI_MASK(
+            "AI_MASK",
+            RowType.of(new DataType[] {DataTypes.STRING()}, new String[] 
{"entities"}),
+            RowType.of(
+                    new DataType[] {DataTypes.STRING(), DataTypes.STRING()},
+                    new String[] {"masked_text", "detected_entities"}),
+            "You are a data masking system. Detect and consistently mask these 
entity types in the "
+                    + "input: %s. Preserve the usefulness and structure of 
non-sensitive content, and "
+                    + "report the entities that were detected.\n");
 
     private final String functionName;
     private final RowType inputType;
@@ -56,6 +117,6 @@ public enum AiTextFunctionDef {
     }
 
     public String buildPrompt(Object... args) {
-        return String.format(promptTemplate, args);
+        return String.format(Locale.ROOT, promptTemplate, args);
     }
 }
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 7d17c86c4f..81aeadc260 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
@@ -33,17 +33,54 @@ import java.util.List;
 /** Built-in AI functions available to transform expressions. */
 public class AiFunctions {
 
+    private static final int MAX_INVALID_JSON_RESPONSE_LENGTH = 512;
+
     private AiFunctions() {}
 
     public static BinaryVariant aiComplete(AiModelClient model, String input, 
String systemPrompt) {
+        return generateText(model, 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 aiTranslate(
+            AiModelClient model, String input, String sourceLang, String 
targetLang) {
+        return generateText(model, 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 aiSentiment(AiModelClient model, String input) 
{
+        return generateText(model, 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 aiMask(AiModelClient model, String input, 
String entities) {
+        return generateText(model, AiTextFunctionDef.AI_MASK, input, entities);
+    }
+
+    private static BinaryVariant generateText(
+            AiModelClient model,
+            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");
         }
 
-        AiTextFunctionDef function = AiTextFunctionDef.AI_COMPLETE;
         String prompt =
-                function.buildPrompt(systemPrompt)
+                function.buildPrompt(promptArguments)
                         + "\n"
                         + buildOutputSchemaHint(function.getOutputType());
         String json = ((SupportsTextGeneration) model).generate(prompt, input);
@@ -53,11 +90,19 @@ public class AiFunctions {
         try {
             return BinaryVariantInternalBuilder.parseJson(json, false);
         } catch (IOException e) {
-            throw new RuntimeException("Failed to parse AI response as JSON: " 
+ json, e);
+            throw new RuntimeException(
+                    "AI function "
+                            + function.getFunctionName()
+                            + " returned invalid JSON: "
+                            + truncateInvalidJsonResponse(json),
+                    e);
         }
     }
 
     public static List<Float> aiEmbed(AiModelClient model, String input) {
+        if (input == null) {
+            return null;
+        }
         if (!(model instanceof SupportsEmbedding)) {
             throw new UnsupportedOperationException(
                     "Model " + model.getClass().getName() + " does not support 
embedding");
@@ -66,8 +111,17 @@ public class AiFunctions {
         return embedding == null ? null : Floats.asList(embedding);
     }
 
+    private static String truncateInvalidJsonResponse(String response) {
+        if (response.length() <= MAX_INVALID_JSON_RESPONSE_LENGTH) {
+            return response;
+        }
+        return response.substring(0, MAX_INVALID_JSON_RESPONSE_LENGTH) + "... 
(truncated)";
+    }
+
     private static String buildOutputSchemaHint(RowType outputType) {
-        StringBuilder builder = new StringBuilder("Return valid JSON with this 
shape:\n{\n");
+        StringBuilder builder =
+                new StringBuilder(
+                        "Return only valid JSON without Markdown fences or 
additional text, using this shape:\n{\n");
         List<String> fieldNames = outputType.getFieldNames();
         for (int i = 0; i < fieldNames.size(); i++) {
             builder.append("  \"")
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 02691968c2..1ee15901bd 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,11 +18,16 @@
 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.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.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;
@@ -901,28 +906,62 @@ public class TransformParser {
     /** 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)) {
-            validateAiModelReferences(parseProjectionExpression(projection), 
declaredModelNames);
+            validateAiModelCapabilities(
+                    parseProjectionExpression(projection), modelClients, 
userDefinedFunctionNames);
         }
         if (!isNullOrWhitespaceOnly(filter)) {
-            validateAiModelReferences(parseFilterExpression(filter), 
declaredModelNames);
+            validateAiModelCapabilities(
+                    parseFilterExpression(filter), modelClients, 
userDefinedFunctionNames);
         }
     }
 
-    private static void validateAiModelReferences(SqlNode node, Set<String> 
declaredModelNames) {
+    private static void validateAiModelReferences(
+            SqlNode node, Set<String> declaredModelNames, Set<String> 
userDefinedFunctionNames) {
         if (node instanceof SqlCall) {
             SqlCall call = (SqlCall) node;
-            if (isAiFunction(call.getOperator().getName())) {
+            String functionName = call.getOperator().getName();
+            if (isAiFunction(functionName)
+                    && !isUserDefinedFunction(functionName, 
userDefinedFunctionNames)) {
                 if (call.operandCount() == 0) {
                     return;
                 }
-                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);
-                String modelName = ((SqlCharStringLiteral) 
modelArgument).getNlsString().getValue();
+                String modelName = resolveAiModelName(call);
                 Preconditions.checkArgument(
                         declaredModelNames.contains(modelName),
                         "Model '%s' referenced by %s has not been declared.",
@@ -931,19 +970,95 @@ public class TransformParser {
             }
             for (SqlNode operand : call.getOperandList()) {
                 if (operand != null) {
-                    validateAiModelReferences(operand, declaredModelNames);
+                    validateAiModelReferences(
+                            operand, declaredModelNames, 
userDefinedFunctionNames);
                 }
             }
         } else if (node instanceof SqlNodeList) {
             for (SqlNode child : (SqlNodeList) node) {
-                validateAiModelReferences(child, declaredModelNames);
+                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);
+                if (isTextAiFunction(functionName)) {
+                    Preconditions.checkArgument(
+                            modelClient instanceof SupportsTextGeneration,
+                            "Model '%s' referenced by %s does not support text 
generation.",
+                            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 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 "AI_COMPLETE".equalsIgnoreCase(functionName)
-                || "AI_EMBED".equalsIgnoreCase(functionName);
+        return isTextAiFunction(functionName) || 
isEmbeddingAiFunction(functionName);
+    }
+
+    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;
     }
 
     public static boolean hasAsterisk(@Nullable String projection) {
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 3cf9873fc8..adf324bf6b 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
@@ -23,6 +23,9 @@ import 
org.apache.flink.cdc.common.model.abilities.SupportsTextGeneration;
 
 import org.junit.jupiter.api.Test;
 
+import java.util.ArrayList;
+import java.util.List;
+
 import static org.assertj.core.api.Assertions.assertThat;
 import static org.assertj.core.api.Assertions.assertThatThrownBy;
 
@@ -34,16 +37,27 @@ class AiFunctionsTest {
 
         private static final long serialVersionUID = 1L;
 
-        private String lastPrompt;
+        private final String response;
+        private final List<String> prompts = new ArrayList<>();
+        private int embedCalls;
+
+        private TestModelClient() {
+            this("{\"result\":\"ABC\"}");
+        }
+
+        private TestModelClient(String response) {
+            this.response = response;
+        }
 
         @Override
         public String generate(String systemPrompt, String userInput) {
-            lastPrompt = systemPrompt;
-            return "{\"result\":\"ABC\"}";
+            prompts.add(systemPrompt);
+            return response;
         }
 
         @Override
         public float[] embed(String text) {
+            embedCalls++;
             return new float[] {0.1f, 0.2f, 0.3f};
         }
     }
@@ -53,13 +67,50 @@ class AiFunctionsTest {
     }
 
     @Test
-    void testAiFunctions() {
+    void testTextAiFunctionsUseEnglishPromptsAndParseJsonResponses() {
         TestModelClient model = new TestModelClient();
 
         assertThat(AiFunctions.aiComplete(model, "input", "Return three 
letters"))
                 .hasToString("{\"result\":\"ABC\"}");
-        assertThat(model.lastPrompt).contains("Return three 
letters").contains("\"result\"");
+        assertThat(AiFunctions.aiClassify(model, "input", "positive,negative"))
+                .hasToString("{\"result\":\"ABC\"}");
+        assertThat(AiFunctions.aiTranslate(model, "input", "auto", "en"))
+                .hasToString("{\"result\":\"ABC\"}");
+        assertThat(AiFunctions.aiSummarize(model, "input", 100))
+                .hasToString("{\"result\":\"ABC\"}");
+        assertThat(AiFunctions.aiSentiment(model, 
"input")).hasToString("{\"result\":\"ABC\"}");
+        assertThat(AiFunctions.aiExtract(model, "input", "name:string"))
+                .hasToString("{\"result\":\"ABC\"}");
+        assertThat(AiFunctions.aiMask(model, "input", "email,phone"))
+                .hasToString("{\"result\":\"ABC\"}");
+
+        assertThat(model.prompts).hasSize(7);
+        assertThat(model.prompts.get(0)).contains("Return three 
letters").contains("\"result\"");
+        assertThat(model.prompts.get(1))
+                .contains("text classifier", "positive,negative", 
"\"category\"");
+        assertThat(model.prompts.get(2))
+                .contains("translator", "auto", "en", "\"translated_text\"");
+        assertThat(model.prompts.get(3))
+                .contains("text summarizer", "100 characters", "\"summary\"");
+        assertThat(model.prompts.get(4))
+                .contains("sentiment analyzer", "\"score\"", "\"confidence\"");
+        assertThat(model.prompts.get(5))
+                .contains("information extraction", "name:string", 
"\"extracted_json\"");
+        assertThat(model.prompts.get(6)).contains("data masking", 
"email,phone", "\"masked_text\"");
+        assertThat(model.prompts)
+                .allSatisfy(
+                        prompt ->
+                                assertThat(prompt)
+                                        .contains("Return only valid JSON")
+                                        .doesNotContainPattern("\\p{IsHan}"));
+    }
+
+    @Test
+    void testEmbeddingFunction() {
+        TestModelClient model = new TestModelClient();
+
         assertThat(AiFunctions.aiEmbed(model, "input")).containsExactly(0.1f, 
0.2f, 0.3f);
+        assertThat(model.embedCalls).isOne();
     }
 
     @Test
@@ -76,16 +127,41 @@ class AiFunctionsTest {
 
     @Test
     void testInvalidJsonResponse() {
-        TestModelClient model =
-                new TestModelClient() {
-                    @Override
-                    public String generate(String systemPrompt, String 
userInput) {
-                        return "not-json";
-                    }
-                };
+        TestModelClient model = new TestModelClient("not-json");
 
-        assertThatThrownBy(() -> AiFunctions.aiComplete(model, "input", 
"prompt"))
+        assertThatThrownBy(() -> AiFunctions.aiClassify(model, "input", 
"positive,negative"))
+                .isInstanceOf(RuntimeException.class)
+                .hasMessage("AI function AI_CLASSIFY returned invalid JSON: 
not-json");
+    }
+
+    @Test
+    void testInvalidJsonResponseIsTruncated() {
+        String longInvalidJson = "x".repeat(600);
+        TestModelClient model = new TestModelClient(longInvalidJson);
+
+        assertThatThrownBy(() -> AiFunctions.aiClassify(model, "input", 
"positive,negative"))
                 .isInstanceOf(RuntimeException.class)
-                .hasMessageContaining("Failed to parse AI response as JSON");
+                .hasMessage(
+                        "AI function AI_CLASSIFY returned invalid JSON: "
+                                + "x".repeat(512)
+                                + "... (truncated)");
+    }
+
+    @Test
+    void testNullInputSkipsModelInvocation() {
+        TestModelClient model = new TestModelClient();
+
+        assertThat(AiFunctions.aiClassify(model, null, 
"positive,negative")).isNull();
+        assertThat(AiFunctions.aiEmbed(model, null)).isNull();
+        assertThat(model.prompts).isEmpty();
+        assertThat(model.embedCalls).isZero();
+    }
+
+    @Test
+    void testNullModelResponseReturnsNull() {
+        TestModelClient model = new TestModelClient(null);
+
+        assertThat(AiFunctions.aiSummarize(model, "input", 100)).isNull();
+        assertThat(model.prompts).hasSize(1);
     }
 }
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 e7af51542e..76cb744a5d 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,15 +17,20 @@
 
 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.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;
 import org.apache.flink.cdc.runtime.operators.transform.ProjectionColumn;
+import 
org.apache.flink.cdc.runtime.operators.transform.UserDefinedFunctionDescriptor;
 
 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;
@@ -40,6 +45,24 @@ class AiFunctionParserTest {
                     Column.physicalColumn("id", DataTypes.INT()),
                     Column.physicalColumn("content", DataTypes.STRING()));
 
+    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];
+        }
+    }
+
     @Test
     void testTranslateAiFunctions() {
         List<ProjectionColumn> columns =
@@ -56,6 +79,69 @@ class AiFunctionParserTest {
                 .containsExactly(DataTypes.VARIANT(), 
DataTypes.ARRAY(DataTypes.FLOAT()));
     }
 
+    @Test
+    void testTranslateSpecializedTextAiFunctions() {
+        List<ProjectionColumn> columns =
+                translate(
+                        "AI_CLASSIFY('model', content, 'positive,negative') AS 
classified, "
+                                + "AI_TRANSLATE('model', content, 'auto', 
'en') AS translated, "
+                                + "AI_SUMMARIZE('model', content, 100) AS 
summarized, "
+                                + "AI_SENTIMENT('model', content) AS 
sentiment, "
+                                + "AI_EXTRACT('model', content, 'name:string') 
AS extracted, "
+                                + "AI_MASK('model', content, 'email,phone') AS 
masked");
+
+        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\")");
+        assertThat(columns)
+                .extracting(ProjectionColumn::getDataType)
+                .containsOnly(DataTypes.VARIANT());
+    }
+
+    @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();
+
+        List<ProjectionColumn> columns =
+                TransformParser.generateProjectionColumns(
+                        "AI_SENTIMENT(id) AS sentiment",
+                        COLUMNS,
+                        List.of(
+                                new UserDefinedFunctionDescriptor(
+                                        "ai_sentiment",
+                                        
"org.apache.flink.cdc.udf.examples.java.AddOneFunctionClass")),
+                        new SupportedMetadataColumn[0]);
+
+        assertThat(columns)
+                .extracting(ProjectionColumn::getScriptExpression)
+                .containsExactly("__udf_ai_sentiment.eval($0)");
+        assertThat(columns)
+                .extracting(ProjectionColumn::getDataType)
+                .containsExactly(DataTypes.STRING());
+    }
+
     @Test
     void testModelArgumentMustBeStringConstant() {
         assertThatThrownBy(
@@ -87,6 +173,53 @@ class AiFunctionParserTest {
                                         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());
+
+        assertThatCode(
+                        () ->
+                                TransformParser.validateAiModelCapabilities(
+                                        "AI_CLASSIFY('textModel', content, 
'a,b') AS classified, "
+                                                + "AI_EMBED('embeddingModel', 
content) AS 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");
     }
 
     @Test
@@ -97,6 +230,18 @@ class AiFunctionParserTest {
                 .hasMessageContaining("Invalid number of arguments to function 
'AI_COMPLETE'");
         assertThatThrownBy(() -> translate("AI_COMPLETE() AS completed"))
                 .hasMessageContaining("Invalid number of arguments to function 
'AI_COMPLETE'");
+        assertThatThrownBy(() -> translate("AI_CLASSIFY('model', content) AS 
classified"))
+                .hasMessageContaining("Invalid number of arguments to function 
'AI_CLASSIFY'");
+        assertThatThrownBy(() -> translate("AI_TRANSLATE('model', content, 
'auto') AS translated"))
+                .hasMessageContaining("Invalid number of arguments to function 
'AI_TRANSLATE'");
+        assertThatThrownBy(() -> translate("AI_SENTIMENT('model') AS 
sentiment"))
+                .hasMessageContaining("Invalid number of arguments to function 
'AI_SENTIMENT'");
+        assertThatThrownBy(() -> translate("AI_EXTRACT('model', content) AS 
extracted"))
+                .hasMessageContaining("Invalid number of arguments to function 
'AI_EXTRACT'");
+        assertThatThrownBy(() -> translate("AI_MASK('model', content) AS 
masked"))
+                .hasMessageContaining("Invalid number of arguments to function 
'AI_MASK'");
+        assertThatThrownBy(() -> translate("AI_SUMMARIZE('model', content, 
TRUE) AS summarized"))
+                .hasMessageContaining("Cannot apply 'AI_SUMMARIZE'");
         assertThatCode(
                         () ->
                                 TransformParser.validateAiModelReferences(

Reply via email to