| ################################################################################ |
| # Licensed to the Apache Software Foundation (ASF) under one |
| # or more contributor license agreements. See the NOTICE file |
| # distributed with this work for additional information |
| # regarding copyright ownership. The ASF licenses this file |
| # to you under the Apache License, Version 2.0 (the |
| # "License"); you may not use this file except in compliance |
| # with the License. You may obtain a copy of the License at |
| # |
| # http://www.apache.org/licenses/LICENSE-2.0 |
| # |
| # Unless required by applicable law or agreed to in writing, software |
| # distributed under the License is distributed on an "AS IS" BASIS, |
| # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. |
| # See the License for the specific language governing permissions and |
| # limitations under the License. |
| ################################################################################# |
| import copy |
| import json |
| import logging |
| import time |
| from typing import TYPE_CHECKING, Dict, List, cast |
| from uuid import UUID |
| |
| from pydantic import BaseModel |
| from pyflink.common import Row |
| from pyflink.common.typeinfo import RowTypeInfo |
| |
| from flink_agents.api.agents.agent import STRUCTURED_OUTPUT |
| from flink_agents.api.agents.react_agent import OutputSchema |
| from flink_agents.api.chat_message import ChatMessage, MessageRole |
| from flink_agents.api.chat_models.java_chat_model import JavaChatModelSetup |
| from flink_agents.api.core_options import ( |
| AgentExecutionOptions, |
| ErrorHandlingStrategy, |
| ) |
| from flink_agents.api.events.chat_event import ChatRequestEvent, ChatResponseEvent |
| from flink_agents.api.events.event import Event |
| from flink_agents.api.events.tool_event import ToolRequestEvent, ToolResponseEvent |
| from flink_agents.api.memory_object import MemoryObject |
| from flink_agents.api.resource import ResourceType |
| from flink_agents.api.runner_context import RunnerContext |
| from flink_agents.plan.actions.action import Action |
| from flink_agents.plan.function import PythonFunction |
| |
| if TYPE_CHECKING: |
| from flink_agents.api.chat_models.chat_model import BaseChatModelSetup |
| |
| _TOOL_CALL_CONTEXT = "_TOOL_CALL_CONTEXT" |
| _TOOL_REQUEST_EVENT_CONTEXT = "_TOOL_REQUEST_EVENT_CONTEXT" |
| _RETRY_STATS_CONTEXT = "_RETRY_STATS_CONTEXT" |
| |
| _logger = logging.getLogger(__name__) |
| |
| |
| # ============================================================================ |
| # Helper Functions for Tool Call Context Management |
| # ============================================================================ |
| def _update_tool_call_context( |
| sensory_memory: MemoryObject, |
| initial_request_id: UUID, |
| initial_messages: List[ChatMessage] | None, |
| added_messages: List[ChatMessage], |
| ) -> List[ChatMessage]: |
| """Append messages to tool call context. |
| |
| The messages maybe chat model response with tool calls, or tool execute results. May |
| initialize the context for initial_request_id if needed. |
| """ |
| # TODO: Because memory doesn't support remove currently, so we use |
| # dict to store tool context in memory and remove the specific |
| # tool context from dict after consuming. This will cause write and |
| # read amplification for we need get the whole dict and overwrite it |
| # to memory each time we update a specific tool context. |
| # After memory supports remove, we can use "TOOL_CALL_CONTEXT/request_id" |
| # to store and remove the specific tool context directly. |
| |
| # init if not exists |
| tool_call_context = sensory_memory.get(_TOOL_CALL_CONTEXT) or {} |
| if initial_request_id not in tool_call_context and initial_messages is not None: |
| tool_call_context[initial_request_id] = copy.deepcopy(initial_messages) |
| |
| tool_call_context[initial_request_id].extend(added_messages) |
| |
| # update tool call context |
| sensory_memory.set(_TOOL_CALL_CONTEXT, tool_call_context) |
| return tool_call_context[initial_request_id] |
| |
| |
| def _save_tool_request_event_context( |
| sensory_memory: MemoryObject, |
| tool_request_event_id: UUID, |
| initial_request_id: UUID, |
| model: str, |
| output_schema: OutputSchema | None, |
| ) -> None: |
| """Save the context for a specific tool request event.""" |
| context = sensory_memory.get(_TOOL_REQUEST_EVENT_CONTEXT) or {} |
| context[tool_request_event_id] = { |
| "initial_request_id": initial_request_id, |
| "model": model, |
| "output_schema": output_schema, |
| } |
| sensory_memory.set(_TOOL_REQUEST_EVENT_CONTEXT, context) |
| |
| |
| def _get_tool_request_event_context( |
| sensory_memory: MemoryObject, request_id: UUID |
| ) -> Dict: |
| """Get and remove the context for a specific tool request event.""" |
| context = sensory_memory.get(_TOOL_REQUEST_EVENT_CONTEXT) or {} |
| removed_context = context.pop(request_id, {}) |
| return removed_context |
| |
| |
| def _accumulate_retry_stats( |
| sensory_memory: MemoryObject, |
| initial_request_id: UUID, |
| retry_count: int, |
| retry_wait_sec: int, |
| ) -> None: |
| """Accumulate retry stats for a given initial request across tool call rounds.""" |
| retry_stats_context = sensory_memory.get(_RETRY_STATS_CONTEXT) or {} |
| stats = retry_stats_context.get(initial_request_id, { |
| "total_retry_count": 0, |
| "total_retry_wait_sec": 0, |
| }) |
| stats["total_retry_count"] += retry_count |
| stats["total_retry_wait_sec"] += retry_wait_sec |
| retry_stats_context[initial_request_id] = stats |
| sensory_memory.set(_RETRY_STATS_CONTEXT, retry_stats_context) |
| |
| |
| def _get_retry_stats( |
| sensory_memory: MemoryObject, |
| initial_request_id: UUID, |
| ) -> dict: |
| """Get accumulated retry stats for a given initial request.""" |
| retry_stats_context = sensory_memory.get(_RETRY_STATS_CONTEXT) or {} |
| return retry_stats_context.get(initial_request_id, { |
| "total_retry_count": 0, |
| "total_retry_wait_sec": 0, |
| }) |
| |
| |
| def _record_retry_metrics( |
| ctx: RunnerContext, model: str, retry_count: int, total_retry_wait_sec: int |
| ) -> None: |
| """Record retry metrics under the connection name if retries occurred.""" |
| if retry_count <= 0: |
| return |
| metric_group = ctx.action_metric_group |
| if metric_group is not None: |
| model_group = metric_group.get_sub_group(model) |
| model_group.get_counter("retryCount").inc(retry_count) |
| model_group.get_counter("retryWaitSec").inc(total_retry_wait_sec) |
| |
| |
| def _handle_tool_calls( |
| response: ChatMessage, |
| initial_request_id: UUID, |
| model: str, |
| messages: List[ChatMessage], |
| output_schema: OutputSchema | None, |
| ctx: RunnerContext, |
| ) -> None: |
| """Handle tool calls in chat response.""" |
| _update_tool_call_context( |
| ctx.sensory_memory, initial_request_id, messages, [response] |
| ) |
| |
| tool_request_event = ToolRequestEvent( |
| model=model, |
| tool_calls=response.tool_calls, |
| ) |
| |
| # save tool request event context |
| _save_tool_request_event_context( |
| ctx.sensory_memory, |
| tool_request_event.id, |
| initial_request_id, |
| model, |
| output_schema, |
| ) |
| |
| ctx.send_event(tool_request_event) |
| |
| |
| def _generate_structured_output( |
| response: ChatMessage, output_schema: OutputSchema |
| ) -> ChatMessage: |
| """Deserialize output to expected output schema.""" |
| output_schema = output_schema.output_schema |
| output = json.loads(response.content.strip()) |
| |
| if isinstance(output_schema, type) and issubclass(output_schema, BaseModel): |
| output = output_schema.model_validate(output) |
| elif isinstance(output_schema, RowTypeInfo): |
| field_names = output_schema.get_field_names() |
| values = {} |
| for field_name in field_names: |
| values[field_name] = output[field_name] |
| output = Row(**values) |
| response.extra_args[STRUCTURED_OUTPUT] = output |
| |
| return response |
| |
| |
| async def chat( |
| initial_request_id: UUID, |
| model: str, |
| messages: List[ChatMessage], |
| output_schema: OutputSchema | None, |
| ctx: RunnerContext, |
| ) -> None: |
| """Chat with llm. |
| |
| If there is no tool call generated, we return the chat response event directly, |
| otherwise, we generate tool request event according to the tool calls in chat model |
| response, and save the request and response messages in tool call context. |
| """ |
| chat_model = cast( |
| "BaseChatModelSetup", ctx.get_resource(model, ResourceType.CHAT_MODEL) |
| ) |
| |
| chat_async = ctx.config.get(AgentExecutionOptions.CHAT_ASYNC) |
| # java chat model doesn't support async execution, |
| # see https://github.com/apache/flink-agents/issues/448 for details. |
| chat_async = chat_async and not isinstance(chat_model, JavaChatModelSetup) |
| |
| error_handling_strategy = ctx.config.get(AgentExecutionOptions.ERROR_HANDLING_STRATEGY) |
| num_retries = 0 |
| retry_wait_interval_sec = 0 |
| if error_handling_strategy == ErrorHandlingStrategy.RETRY: |
| num_retries = max(0, ctx.config.get(AgentExecutionOptions.MAX_RETRIES)) |
| retry_wait_interval_config = ctx.config.get( |
| AgentExecutionOptions.RETRY_WAIT_INTERVAL |
| ) |
| retry_wait_interval_sec = ( |
| max(0, retry_wait_interval_config) if retry_wait_interval_config else 0 |
| ) |
| |
| response = None |
| actual_retry_count = 0 |
| total_wait_time_sec = 0 |
| |
| for attempt in range(num_retries + 1): |
| try: |
| if chat_async: |
| response = await ctx.durable_execute_async(chat_model.chat, messages) |
| else: |
| response = ctx.durable_execute(chat_model.chat, messages) |
| |
| if response.extra_args.get("model_name") and response.extra_args.get("promptTokens") and response.extra_args.get("completionTokens"): |
| chat_model._record_token_metrics(response.extra_args["model_name"], response.extra_args["promptTokens"], response.extra_args["completionTokens"]) |
| if output_schema is not None and len(response.tool_calls) == 0: |
| response = _generate_structured_output(response, output_schema) |
| break |
| except Exception as e: |
| if error_handling_strategy == ErrorHandlingStrategy.IGNORE: |
| _logger.warning( |
| f"Chat request {initial_request_id} failed with error: {e}, ignored." |
| ) |
| return |
| elif error_handling_strategy == ErrorHandlingStrategy.RETRY: |
| if attempt == num_retries: |
| raise |
| actual_retry_count = attempt + 1 |
| current_wait_sec = retry_wait_interval_sec * ( |
| 1 << (actual_retry_count - 1) |
| ) |
| _logger.warning( |
| f"Chat request {initial_request_id} failed with error: {e}, " |
| f"retrying {actual_retry_count} / {num_retries}, " |
| f"waiting {current_wait_sec} s." |
| ) |
| if current_wait_sec > 0: |
| time.sleep(current_wait_sec) |
| total_wait_time_sec += current_wait_sec |
| else: |
| _logger.debug( |
| f"Chat request {initial_request_id} failed, the input chat messages are {messages}." |
| ) |
| raise |
| |
| if actual_retry_count > 0: |
| _accumulate_retry_stats( |
| ctx.sensory_memory, initial_request_id, actual_retry_count, total_wait_time_sec |
| ) |
| |
| if ( |
| len(response.tool_calls) > 0 |
| ): # generate tool request event according tool calls in response |
| _handle_tool_calls( |
| response, initial_request_id, model, messages, output_schema, ctx |
| ) |
| else: # if there is no tool call generated, return chat response directly |
| retry_stats = _get_retry_stats(ctx.sensory_memory, initial_request_id) |
| total_retry_count = retry_stats["total_retry_count"] |
| total_retry_wait_sec = retry_stats["total_retry_wait_sec"] |
| |
| _record_retry_metrics(ctx, chat_model.connection, total_retry_count, total_retry_wait_sec) |
| |
| ctx.send_event( |
| ChatResponseEvent( |
| request_id=initial_request_id, |
| response=response, |
| retry_count=total_retry_count, |
| total_retry_wait_sec=total_retry_wait_sec, |
| ) |
| ) |
| |
| |
| async def _process_chat_request(event: ChatRequestEvent, ctx: RunnerContext) -> None: |
| """Process chat request event.""" |
| await chat( |
| initial_request_id=event.id, |
| model=event.model, |
| messages=event.messages, |
| output_schema=event.output_schema, |
| ctx=ctx, |
| ) |
| |
| |
| async def _process_tool_response(event: ToolResponseEvent, ctx: RunnerContext) -> None: |
| """Organize the tool call context and return it to the LLM.""" |
| sensory_memory = ctx.sensory_memory |
| request_id = event.request_id |
| |
| # get correspond tool request event context |
| tool_request_event_context = _get_tool_request_event_context( |
| sensory_memory, request_id |
| ) |
| initial_request_id = tool_request_event_context["initial_request_id"] |
| |
| # update tool call context, and get the entire chat messages. |
| messages = _update_tool_call_context( |
| sensory_memory, |
| initial_request_id, |
| None, |
| [ |
| ChatMessage( |
| role=MessageRole.TOOL, |
| content=str(response), |
| extra_args={"external_id": event.external_ids.get(tool_id)} |
| if event.external_ids and event.external_ids.get(tool_id) |
| else {}, |
| ) |
| for tool_id, response in event.responses.items() |
| ], |
| ) |
| |
| await chat( |
| initial_request_id=initial_request_id, |
| model=tool_request_event_context["model"], |
| messages=messages, |
| output_schema=tool_request_event_context["output_schema"], |
| ctx=ctx, |
| ) |
| |
| |
| async def process_chat_request_or_tool_response( |
| event: Event, ctx: RunnerContext |
| ) -> None: |
| """Built-in action for processing a chat request or tool response. |
| |
| This action listens to ChatRequestEvent and ToolResponseEvent, and handles |
| the complete chat flow including tool calls. It uses sensory memory to save |
| the tool call context, which is a dict mapping request id to chat messages. |
| """ |
| # To avoid https://github.com/alibaba/pemja/issues/88, we log a message here. |
| logging.debug("Processing chat request asynchronously.") |
| if isinstance(event, ChatRequestEvent): |
| await _process_chat_request(event, ctx) |
| elif isinstance(event, ToolResponseEvent): |
| await _process_tool_response(event, ctx) |
| |
| |
| CHAT_MODEL_ACTION = Action( |
| name="chat_model_action", |
| exec=PythonFunction.from_callable(process_chat_request_or_tool_response), |
| listen_event_types=[ |
| f"{ChatRequestEvent.__module__}.{ChatRequestEvent.__name__}", |
| f"{ToolResponseEvent.__module__}.{ToolResponseEvent.__name__}", |
| ], |
| ) |