diff --git a/aidial_assistant/application/assistant_application.py b/aidial_assistant/application/assistant_application.py index ed15f0c..4abc2f2 100644 --- a/aidial_assistant/application/assistant_application.py +++ b/aidial_assistant/application/assistant_application.py @@ -1,14 +1,26 @@ import logging from pathlib import Path -from typing import Tuple +from typing import Callable, Tuple from aidial_sdk.chat_completion import FinishReason from aidial_sdk.chat_completion.base import ChatCompletion -from aidial_sdk.chat_completion.request import Addon, Message, Request, Role +from aidial_sdk.chat_completion.request import Request from aidial_sdk.chat_completion.response import Response +from aidial_sdk.deployment.tokenize import ( + TokenizeError, + TokenizeRequest, + TokenizeResponse, + TokenizeSuccess, +) +from aidial_sdk.deployment.truncate_prompt import ( + TruncatePromptError, + TruncatePromptRequest, + TruncatePromptResponse, + TruncatePromptSuccess, +) from openai.lib.azure import AsyncAzureOpenAI from openai.types.chat import ChatCompletionToolParam -from pydantic import BaseModel +from typing_extensions import override from aidial_assistant.application.addons_dialogue_limiter import ( AddonsDialogueLimiter, @@ -21,19 +33,33 @@ MAIN_BEST_EFFORT_TEMPLATE, MAIN_SYSTEM_DIALOG_MESSAGE, ) +from aidial_assistant.application.request_data import ( + PluginInfo, + RequestData, + get_discarded_user_messages, +) from aidial_assistant.chain.command_chain import ( CommandChain, CommandConstructor, CommandDict, ) -from aidial_assistant.chain.history import History +from aidial_assistant.chain.history import History, ScopedMessage from aidial_assistant.commands.reply import Reply -from aidial_assistant.commands.run_plugin import PluginInfo, RunPlugin +from aidial_assistant.commands.run_plugin import RunPlugin from aidial_assistant.commands.run_tool import RunTool from aidial_assistant.model.model_client import ( ModelClient, + ModelClientRequest, ReasonLengthException, ) +from aidial_assistant.model.tokenize_client import ( + TokenizeClient, + TokenizeClientRequest, +) +from aidial_assistant.model.truncate_propmt_client import ( + TruncatePromptClient, + TruncatePromptClientRequest, +) from aidial_assistant.tools_chain.tools_chain import ( CommandToolDict, ToolsChain, @@ -41,71 +67,87 @@ ) from aidial_assistant.utils.exceptions import ( RequestParameterValidationError, + UnauthorizedAddonError, unhandled_exception_handler, ) from aidial_assistant.utils.open_ai import construct_tool -from aidial_assistant.utils.open_ai_plugin import ( - AddonTokenSource, - get_open_ai_plugin_info, - get_plugin_auth, -) -from aidial_assistant.utils.state import State, parse_history +from aidial_assistant.utils.open_ai_plugin import AddonTokenSource, AIPluginConf +from aidial_assistant.utils.state import State logger = logging.getLogger(__name__) -class AddonReference(BaseModel): - name: str | None - url: str +def _construct_tool(plugin_conf: AIPluginConf) -> ChatCompletionToolParam: + return construct_tool( + plugin_conf.name_for_model, + plugin_conf.description_for_human, + { + "query": { + "type": "string", + "description": "A task written in natural language", + } + }, + ["query"], + ) -def _get_request_args(request: Request) -> dict[str, str]: - args = { - "model": request.model, - "temperature": request.temperature, - "user": request.user, +def _create_history( + messages: list[ScopedMessage], plugins: list[PluginInfo] +) -> History: + plugin_descriptions = { + plugin.info.ai_plugin.name_for_model: plugin.info.open_api.info.description + or plugin.info.ai_plugin.description_for_human + for plugin in plugins } + return History( + assistant_system_message_template=MAIN_SYSTEM_DIALOG_MESSAGE.build( + addons=plugin_descriptions + ), + best_effort_template=MAIN_BEST_EFFORT_TEMPLATE.build( + addons=plugin_descriptions + ), + scoped_messages=messages, + ) - return {k: v for k, v in args.items() if v is not None} +def _get_plugin_auth( + plugin: PluginInfo, token_source: AddonTokenSource +) -> str | None: + auth_type = plugin.info.ai_plugin.auth.type -def _validate_addons(addons: list[Addon] | None) -> list[AddonReference]: - addon_references: list[AddonReference] = [] - for index, addon in enumerate(addons or []): - if addon.url is None: - raise RequestParameterValidationError( - f"Missing required addon url at index {index}.", - param="addons", - ) + if auth_type == "none": + return token_source.default_auth - addon_references.append(AddonReference(name=addon.name, url=addon.url)) + if auth_type == "service_http": + service_token = token_source.get_token(plugin.url) + if service_token is None: + raise UnauthorizedAddonError(f"Missing token for {plugin.url}") - return addon_references + authorization_type = plugin.info.ai_plugin.auth.authorization_type + # Capitalizing because Wolfram, for instance, doesn't like lowercase bearer + return f"{authorization_type.capitalize()} {service_token}" -def _validate_messages(messages: list[Message]) -> None: - if not messages: - raise RequestParameterValidationError( - "Message list cannot be empty.", param="messages" - ) + raise UnauthorizedAddonError(f"Unknown auth type {auth_type}") - if messages[-1].role != Role.USER: - raise RequestParameterValidationError( - "Last message must be from the user.", param="messages" - ) +def _native_tools_request(request_data: RequestData) -> ModelClientRequest: + tools = [ + _construct_tool(plugin.info.ai_plugin) + for plugin in request_data.plugins + ] + return ModelClientRequest( + messages=convert_commands_to_tools(request_data.messages), + max_prompt_tokens=request_data.max_prompt_tokens, + tools=tools, + ) -def _construct_tool(name: str, description: str) -> ChatCompletionToolParam: - return construct_tool( - name, - description, - { - "query": { - "type": "string", - "description": "A task written in natural language", - } - }, - ["query"], + +def _emulated_tools_request(request_data: RequestData) -> ModelClientRequest: + history = _create_history(request_data.messages, request_data.plugins) + return ModelClientRequest( + messages=history.to_protocol_messages(), + max_prompt_tokens=request_data.max_prompt_tokens, ) @@ -116,13 +158,12 @@ def __init__( self.args = parse_args(config_dir) self.tools_supporting_deployments = tools_supporting_deployments + @override @unhandled_exception_handler async def chat_completion( self, request: Request, response: Response ) -> None: - _validate_messages(request.messages) - addon_references = _validate_addons(request.addons) - chat_args = _get_request_args(request) + request_data = await RequestData.from_dial_request(request) model = ModelClient( client=AsyncAzureOpenAI( @@ -131,66 +172,43 @@ async def chat_completion( # 2023-12-01-preview is needed to support tools api_version="2023-12-01-preview", ), - model_args=chat_args, + model_args=request_data.model_args, ) - token_source = AddonTokenSource( request.headers, - (addon_reference.url for addon_reference in addon_references), + (plugin.url for plugin in request_data.plugins), ) - plugins: list[PluginInfo] = [] - # DIAL Core has own names for addons, so in stages we need to map them to the names used by the user - addon_name_mapping: dict[str, str] = {} - for addon_reference in addon_references: - info = await get_open_ai_plugin_info(addon_reference.url) - plugins.append( - PluginInfo( - info=info, - auth=get_plugin_auth( - info.ai_plugin.auth.type, - info.ai_plugin.auth.authorization_type, - addon_reference.url, - token_source, - ), - ) - ) - - if addon_reference.name: - addon_name_mapping[ - info.ai_plugin.name_for_model - ] = addon_reference.name - - if request.model in self.tools_supporting_deployments: + if self._supports_native_tools(request_data.model): await AssistantApplication._run_native_tools_chat( - model, plugins, addon_name_mapping, request, response + model, token_source, request_data, response ) else: await AssistantApplication._run_emulated_tools_chat( - model, plugins, addon_name_mapping, request, response + model, token_source, request_data, response ) @staticmethod async def _run_emulated_tools_chat( model: ModelClient, - addons: list[PluginInfo], - addon_name_mapping: dict[str, str], - request: Request, + token_source: AddonTokenSource, + request_data: RequestData, response: Response, ): - # TODO: Add max_addons_dialogue_tokens as a request parameter - max_addons_dialogue_tokens = 1000 - - def create_command(addon: PluginInfo): - return lambda: RunPlugin(model, addon, max_addons_dialogue_tokens) + def create_command(plugin: PluginInfo, auth: str | None): + return lambda: RunPlugin( + model, plugin, auth, request_data.max_addons_dialogue_tokens + ) command_dict: CommandDict = { - addon.info.ai_plugin.name_for_model: create_command(addon) - for addon in addons + plugin.info.ai_plugin.name_for_model: create_command( + plugin, _get_plugin_auth(plugin, token_source) + ) + for plugin in request_data.plugins } if Reply.token() in command_dict: RequestParameterValidationError( - f"Addon with name '{Reply.token()}' is not allowed for model {request.model}.", + f"Addon with name '{Reply.token()}' is not allowed in emulated tools mode.", param="addons", ) @@ -199,36 +217,27 @@ def create_command(addon: PluginInfo): chain = CommandChain( model_client=model, name="ASSISTANT", command_dict=command_dict ) - addon_descriptions = { - addon.info.ai_plugin.name_for_model: addon.info.open_api.info.description - or addon.info.ai_plugin.description_for_human - for addon in addons - } - history = History( - assistant_system_message_template=MAIN_SYSTEM_DIALOG_MESSAGE.build( - addons=addon_descriptions - ), - best_effort_template=MAIN_BEST_EFFORT_TEMPLATE.build( - addons=addon_descriptions - ), - scoped_messages=parse_history(request.messages), - ) - discarded_messages: int | None = None - if request.max_prompt_tokens is not None: - original_size = history.user_message_count - history = await history.truncate(request.max_prompt_tokens, model) - truncated_size = history.user_message_count - discarded_messages = original_size - truncated_size + history = _create_history(request_data.messages, request_data.plugins) + discarded_user_messages: list[int] | None = None + if request_data.max_prompt_tokens is not None: + history, discarded_messages = await history.truncate( + model, request_data.max_prompt_tokens + ) + discarded_user_messages = get_discarded_user_messages( + request_data.messages, discarded_messages + ) # TODO: else compare the history size to the max prompt tokens of the underlying model choice = response.create_single_choice() choice.open() - callback = AssistantChainCallback(choice, addon_name_mapping) + callback = AssistantChainCallback( + choice, request_data.addon_name_mapping + ) finish_reason = FinishReason.STOP try: model_request_limiter = AddonsDialogueLimiter( - max_addons_dialogue_tokens, model + request_data.max_addons_dialogue_tokens, model ) await chain.run_chat(history, callback, model_request_limiter) except ReasonLengthException: @@ -243,45 +252,42 @@ def create_command(addon: PluginInfo): model.total_prompt_tokens, model.total_completion_tokens ) - if discarded_messages is not None: - response.set_discarded_messages(discarded_messages) + if discarded_user_messages is not None: + response.set_discarded_messages(discarded_user_messages) @staticmethod async def _run_native_tools_chat( model: ModelClient, - plugins: list[PluginInfo], - addon_name_mapping: dict[str, str], - request: Request, + token_source: AddonTokenSource, + request_data: RequestData, response: Response, ): - # TODO: Add max_addons_dialogue_tokens as a request parameter - max_addons_dialogue_tokens = 1000 - def create_command_tool( - plugin: PluginInfo, + plugin: PluginInfo, auth: str | None ) -> Tuple[CommandConstructor, ChatCompletionToolParam]: return lambda: RunTool( - model, plugin, max_addons_dialogue_tokens - ), _construct_tool( - plugin.info.ai_plugin.name_for_model, - plugin.info.ai_plugin.description_for_human, - ) + model, plugin, auth, request_data.max_addons_dialogue_tokens + ), _construct_tool(plugin.info.ai_plugin) commands: CommandToolDict = { - plugin.info.ai_plugin.name_for_model: create_command_tool(plugin) - for plugin in plugins + plugin.info.ai_plugin.name_for_model: create_command_tool( + plugin, _get_plugin_auth(plugin, token_source) + ) + for plugin in request_data.plugins } chain = ToolsChain(model, commands) choice = response.create_single_choice() choice.open() - callback = AssistantChainCallback(choice, addon_name_mapping) + callback = AssistantChainCallback( + choice, request_data.addon_name_mapping + ) finish_reason = FinishReason.STOP - messages = convert_commands_to_tools(parse_history(request.messages)) + messages = convert_commands_to_tools(request_data.messages) try: model_request_limiter = AddonsDialogueLimiter( - max_addons_dialogue_tokens, model + request_data.max_addons_dialogue_tokens, model ) await chain.run_chat(messages, callback, model_request_limiter) except ReasonLengthException: @@ -294,3 +300,73 @@ def create_command_tool( response.set_usage( model.total_prompt_tokens, model.total_completion_tokens ) + + @override + async def tokenize(self, request: TokenizeRequest) -> TokenizeResponse: + inputs: list[ModelClientRequest | str] = [] + + for tokenizer_input in request.inputs: + if tokenizer_input.type == "string": + inputs.append(tokenizer_input.value) + continue + + request_data = await RequestData.from_dial_request( + tokenizer_input.value + ) + if self._supports_native_tools(request_data.model): + inputs.append(_native_tools_request(request_data)) + else: + inputs.append(_emulated_tools_request(request_data)) + + client = TokenizeClient(self.args.openai_conf.api_base) + outputs = await client.tokenize(TokenizeClientRequest(inputs=inputs)) + + return TokenizeResponse( + outputs=[ + TokenizeSuccess(token_count=output) + if isinstance(output, int) + else TokenizeError(error=output) + for output in outputs + ] + ) + + @override + async def truncate_prompt( + self, request: TruncatePromptRequest + ) -> TruncatePromptResponse: + inputs: list[ModelClientRequest] = [] + indices_converters: list[Callable[[list[int]], list[int]]] = [] + + for completion_request in request.inputs: + request_data = await RequestData.from_dial_request( + completion_request + ) + if self._supports_native_tools(request_data.model): + inputs.append(_native_tools_request(request_data)) + indices_converters.append(lambda indices: indices) + else: + inputs.append(_emulated_tools_request(request_data)) + indices_converters.append( + lambda indices: get_discarded_user_messages( + request_data.messages, indices + ) + ) + + client = TruncatePromptClient(self.args.openai_conf.api_base) + outputs = await client.truncate_prompt( + TruncatePromptClientRequest(inputs=inputs) + ) + + return TruncatePromptResponse( + outputs=[ + TruncatePromptSuccess( + discarded_messages=indices_converters[index](output) + ) + if isinstance(output, list) + else TruncatePromptError(error=output) + for index, output in enumerate(outputs) + ] + ) + + def _supports_native_tools(self, model: str) -> bool: + return model in self.tools_supporting_deployments diff --git a/aidial_assistant/application/request_data.py b/aidial_assistant/application/request_data.py new file mode 100644 index 0000000..08c139a --- /dev/null +++ b/aidial_assistant/application/request_data.py @@ -0,0 +1,228 @@ +import json +from enum import Enum + +from aidial_sdk.chat_completion import Addon, CustomContent, Message, Role +from aidial_sdk.chat_completion.request import ChatCompletionRequest +from openai.types.chat import ChatCompletionMessageParam +from pydantic import BaseModel + +from aidial_assistant.chain.command_result import ( + CommandInvocation, + commands_to_text, +) +from aidial_assistant.utils.exceptions import RequestParameterValidationError +from aidial_assistant.utils.open_ai import ( + assistant_message, + system_message, + user_message, +) +from aidial_assistant.utils.open_ai_plugin import ( + OpenAIPluginInfo, + get_open_ai_plugin_info, +) +from aidial_assistant.utils.state import Invocation, State + + +class AddonReference(BaseModel): + name: str | None + url: str + + +class PluginInfo(BaseModel): + info: OpenAIPluginInfo + url: str + + +class MessageScope(str, Enum): + INTERNAL = "internal" # internal dialog with plugins/addons, not visible to the user on the top level + USER = "user" # top-level dialog with the user + + +class ScopedMessage(BaseModel): + scope: MessageScope = MessageScope.USER + message: ChatCompletionMessageParam + user_index: int + + +def _validate_required(value: str | None, name: str) -> str: + if value is None: + raise RequestParameterValidationError( + f"Missing required parameter {name}.", param=name + ) + + return value + + +def _validate_messages(messages: list[Message]) -> None: + if not messages: + raise RequestParameterValidationError( + "Message list cannot be empty.", param="messages" + ) + + if messages[-1].role != Role.USER: + raise RequestParameterValidationError( + "Last message must be from the user.", param="messages" + ) + + +def _validate_addons(addons: list[Addon] | None) -> list[AddonReference]: + addon_references: list[AddonReference] = [] + for index, addon in enumerate(addons or []): + if addon.url is None: + raise RequestParameterValidationError( + f"Missing required addon url at index {index}.", + param="addons", + ) + + addon_references.append(AddonReference(name=addon.name, url=addon.url)) + + return addon_references + + +def _get_model_args(request: ChatCompletionRequest) -> dict[str, str]: + args = { + "model": request.model, + "temperature": request.temperature, + "user": request.user, + } + + return {k: v for k, v in args.items() if v is not None} + + +def _convert_old_commands(string: str) -> str: + """Converts old commands to new format. + Previously saved conversations with assistant will stop working if state is not updated. + + Old format: + {"commands": [{"command": "run-addon", "args": ["", ""]}]} + New format: + {"commands": [{"command": "", "arguments": {"query": ""}}]} + """ + commands = json.loads(string) + result: list[CommandInvocation] = [] + + for command in commands["commands"]: + command_name = command["command"] + # run-addon was previously called run-plugin + if command_name in ("run-addon", "run-plugin"): + args = command["args"] + result.append( + CommandInvocation(command=args[0], arguments={"query": args[1]}) + ) + else: + result.append(command) + + return commands_to_text(result) + + +def _get_invocations(custom_content: CustomContent | None) -> list[Invocation]: + if custom_content is None: + return [] + + state: State | None = custom_content.state + if state is None: + return [] + + invocations: list[Invocation] | None = state.get("invocations") + if invocations is None: + return [] + + invocations.sort(key=lambda invocation: int(invocation["index"])) + return invocations + + +def _parse_history(history: list[Message]) -> list[ScopedMessage]: + messages: list[ScopedMessage] = [] + for index, message in enumerate(history): + if message.role == Role.ASSISTANT: + invocations = _get_invocations(message.custom_content) + for invocation in invocations: + messages.append( + ScopedMessage( + scope=MessageScope.INTERNAL, + message=assistant_message( + _convert_old_commands(invocation["request"]) + ), + user_index=index, + ) + ) + messages.append( + ScopedMessage( + scope=MessageScope.INTERNAL, + message=user_message(invocation["response"]), + user_index=index, + ) + ) + + messages.append( + ScopedMessage( + message=assistant_message(message.content or ""), + user_index=index, + ) + ) + elif message.role == Role.USER: + messages.append( + ScopedMessage( + message=user_message(message.content or ""), + user_index=index, + ) + ) + elif message.role == Role.SYSTEM: + messages.append( + ScopedMessage( + message=system_message(message.content or ""), + user_index=index, + ) + ) + else: + raise RequestParameterValidationError( + f"Role {message.role} is not supported.", param="messages" + ) + + return messages + + +def get_discarded_user_messages( + scoped_messages: list[ScopedMessage], discarded_messages: list[int] +) -> list[int]: + return [scoped_messages[index].user_index for index in discarded_messages] + + +class RequestData(BaseModel): + model: str + model_args: dict[str, str] + messages: list[ScopedMessage] + plugins: list[PluginInfo] + addon_name_mapping: dict[str, str] + max_prompt_tokens: int | None = None + # TODO: Add max_addons_dialogue_tokens as a request parameter to the dial sdk + max_addons_dialogue_tokens: int = 1000 + + @classmethod + async def from_dial_request( + cls, request: ChatCompletionRequest + ) -> "RequestData": + _validate_messages(request.messages) + addon_references = _validate_addons(request.addons) + model = _validate_required(request.model, "model") + + plugins: list[PluginInfo] = [] + # DIAL Core has own names for addons, so in stages we need to map them to the names used by the user + addon_name_mapping: dict[str, str] = {} + for addon_reference in addon_references: + info = await get_open_ai_plugin_info(addon_reference.url) + plugins.append(PluginInfo(info=info, url=addon_reference.url)) + + if addon_reference.name: + addon_name_mapping[ + info.ai_plugin.name_for_model + ] = addon_reference.name + + return cls( + model=model, + model_args=_get_model_args(request), + messages=_parse_history(request.messages), + plugins=plugins, + addon_name_mapping=addon_name_mapping, + max_prompt_tokens=request.max_prompt_tokens, + ) diff --git a/aidial_assistant/chain/command_chain.py b/aidial_assistant/chain/command_chain.py index ea76ab4..2c8aa72 100644 --- a/aidial_assistant/chain/command_chain.py +++ b/aidial_assistant/chain/command_chain.py @@ -36,6 +36,7 @@ from aidial_assistant.model.model_client import ( ChatCompletionMessageParam, ModelClient, + ModelClientRequest, ) from aidial_assistant.utils.stream import CumulativeStream @@ -72,11 +73,7 @@ def __init__( self.name = name self.model_client = model_client self.command_dict = command_dict - self.model_extra_args = ( - {} - if max_completion_tokens is None - else {"max_tokens": max_completion_tokens} - ) + self.max_completion_tokens = max_completion_tokens self.max_retry_count = max_retry_count def _log_message(self, role: str, content: str | None): @@ -149,7 +146,10 @@ async def _run_with_protocol_failure_retries( chunk_stream = CumulativeStream( self.model_client.agenerate( - all_messages, **self.model_extra_args # type: ignore + ModelClientRequest( + messages=all_messages, + max_tokens=self.max_completion_tokens, + ) ) ) try: @@ -252,7 +252,9 @@ async def _generate_result( messages: list[ChatCompletionMessageParam], callback: ChainCallback, ): - stream = self.model_client.agenerate(messages) + stream = self.model_client.agenerate( + ModelClientRequest(messages=messages) + ) await CommandChain._to_result(stream, callback.result_callback()) diff --git a/aidial_assistant/chain/history.py b/aidial_assistant/chain/history.py index 6e8db05..a53e07c 100644 --- a/aidial_assistant/chain/history.py +++ b/aidial_assistant/chain/history.py @@ -1,9 +1,15 @@ -from enum import Enum +from typing import Tuple, cast from jinja2 import Template -from openai.types.chat import ChatCompletionMessageParam -from pydantic import BaseModel +from openai.types.chat import ( + ChatCompletionMessageParam, + ChatCompletionSystemMessageParam, +) +from aidial_assistant.application.request_data import ( + MessageScope, + ScopedMessage, +) from aidial_assistant.chain.command_result import ( CommandInvocation, commands_to_text, @@ -18,16 +24,6 @@ class ContextLengthExceeded(Exception): pass -class MessageScope(str, Enum): - INTERNAL = "internal" # internal dialog with plugins/addons, not visible to the user on the top level - USER = "user" # top-level dialog with the user - - -class ScopedMessage(BaseModel): - scope: MessageScope = MessageScope.USER - message: ChatCompletionMessageParam - - class History: def __init__( self, @@ -40,35 +36,32 @@ def __init__( ) self.best_effort_template = best_effort_template self.scoped_messages = scoped_messages - self._user_message_count = sum( - 1 - for message in scoped_messages - if message.scope == MessageScope.USER - ) def to_protocol_messages(self) -> list[ChatCompletionMessageParam]: messages: list[ChatCompletionMessageParam] = [] - for index, scoped_message in enumerate(self.scoped_messages): + scoped_message_iterator = iter(self.scoped_messages) + if self._is_first_system_message(): + message = cast( + ChatCompletionSystemMessageParam, + next(scoped_message_iterator).message, + ) + messages.append( + system_message( + self.assistant_system_message_template.render( + system_prefix=message["content"] + ) + ) + ) + else: + messages.append( + system_message(self.assistant_system_message_template.render()) + ) + + for scoped_message in scoped_message_iterator: message = scoped_message.message scope = scoped_message.scope - if index == 0: - if message["role"] == "system": - messages.append( - system_message( - self.assistant_system_message_template.render( - system_prefix=message["content"] - ) - ) - ) - else: - messages.append( - system_message( - self.assistant_system_message_template.render() - ) - ) - messages.append(message) - elif scope == MessageScope.USER and message["role"] == "assistant": + if scope == MessageScope.USER and message["role"] == "assistant": # Clients see replies in plain text, but the model should understand how to reply appropriately. content = commands_to_text( [ @@ -107,51 +100,59 @@ def to_best_effort_messages( return messages async def truncate( - self, max_prompt_tokens: int, model_client: ModelClient - ) -> "History": - discarded_messages = await model_client.get_discarded_messages( - self.to_protocol_messages(), - max_prompt_tokens, + self, model_client: ModelClient, max_prompt_tokens: int + ) -> Tuple["History", list[int]]: + discarded_messages = await self._get_discarded_messages( + model_client, max_prompt_tokens ) - if discarded_messages > 0: - return History( + if not discarded_messages: + return self, [] + + discarded_messages_set = set(discarded_messages) + return ( + History( assistant_system_message_template=self.assistant_system_message_template, best_effort_template=self.best_effort_template, - scoped_messages=self._skip_messages(discarded_messages), - ) - - return self - - @property - def user_message_count(self) -> int: - return self._user_message_count - - def _skip_messages(self, discarded_messages: int) -> list[ScopedMessage]: - messages: list[ScopedMessage] = [] - current_message = self.scoped_messages[0] - message_iterator = iter(self.scoped_messages) - for _ in range(discarded_messages): - current_message = next(message_iterator) - while current_message.message["role"] == "system": - # System messages should be kept in the history - messages.append(current_message) - current_message = next(message_iterator) + scoped_messages=[ + scoped_message + for index, scoped_message in enumerate(self.scoped_messages) + if index not in discarded_messages_set + ], + ), + discarded_messages, + ) - if current_message.scope == MessageScope.INTERNAL: - while current_message.scope == MessageScope.INTERNAL: - current_message = next(message_iterator) + async def _get_discarded_messages( + self, model_client: ModelClient, max_prompt_tokens: int + ) -> list[int]: + discarded_protocol_messages = await model_client.get_discarded_messages( + self.to_protocol_messages(), + max_prompt_tokens, + ) - # Internal messages (i.e. addon requests/responses) are always followed by an assistant reply - assert ( - current_message.message["role"] == "assistant" - ), "Internal messages must be followed by an assistant reply." + if discarded_protocol_messages: + discarded_protocol_messages.sort() + discarded_messages = ( + discarded_protocol_messages + if self._is_first_system_message() + else [index - 1 for index in discarded_protocol_messages] + ) + user_indices = set( + self.scoped_messages[index].user_index + for index in discarded_messages + ) - remaining_messages = list(message_iterator) - assert ( - len(remaining_messages) > 0 - ), "No user messages left after history truncation." + return [ + index + for index, scoped_message in enumerate(self.scoped_messages) + if scoped_message.user_index in user_indices + ] - messages += remaining_messages + return discarded_protocol_messages - return messages + def _is_first_system_message(self) -> bool: + return ( + len(self.scoped_messages) > 0 + and self.scoped_messages[0].message["role"] == "system" + ) diff --git a/aidial_assistant/commands/run_plugin.py b/aidial_assistant/commands/run_plugin.py index 96f2913..9c9d6d9 100644 --- a/aidial_assistant/commands/run_plugin.py +++ b/aidial_assistant/commands/run_plugin.py @@ -1,11 +1,11 @@ from langchain.tools import APIOperation -from pydantic.main import BaseModel from typing_extensions import override from aidial_assistant.application.prompts import ( ADDON_BEST_EFFORT_TEMPLATE, ADDON_SYSTEM_DIALOG_MESSAGE, ) +from aidial_assistant.application.request_data import PluginInfo from aidial_assistant.chain.command_chain import ( CommandChain, CommandConstructor, @@ -27,12 +27,6 @@ ) from aidial_assistant.open_api.operation_selector import collect_operations from aidial_assistant.utils.open_ai import user_message -from aidial_assistant.utils.open_ai_plugin import OpenAIPluginInfo - - -class PluginInfo(BaseModel): - info: OpenAIPluginInfo - auth: str | None class RunPlugin(Command): @@ -40,10 +34,12 @@ def __init__( self, model_client: ModelClient, plugin: PluginInfo, + auth: str | None, max_completion_tokens: int, ): self.model_client = model_client self.plugin = plugin + self.auth = auth self.max_completion_tokens = max_completion_tokens @staticmethod @@ -66,7 +62,7 @@ async def _run_plugin( api_schema = "\n\n".join([op.to_typescript() for op in ops.values()]) # type: ignore def create_command(op: APIOperation): - return lambda: OpenAPIChatCommand(op, self.plugin.auth) + return lambda: OpenAPIChatCommand(op, self.auth) command_dict: dict[str, CommandConstructor] = {} for name, op in ops.items(): @@ -87,7 +83,9 @@ def create_command(op: APIOperation): best_effort_template=ADDON_BEST_EFFORT_TEMPLATE.build( api_schema=api_schema ), - scoped_messages=[ScopedMessage(message=user_message(query))], + scoped_messages=[ + ScopedMessage(message=user_message(query), user_index=0) + ], ) chat = CommandChain( diff --git a/aidial_assistant/commands/run_tool.py b/aidial_assistant/commands/run_tool.py index 5c6b4a2..0723908 100644 --- a/aidial_assistant/commands/run_tool.py +++ b/aidial_assistant/commands/run_tool.py @@ -7,6 +7,7 @@ from openai.types.chat import ChatCompletionToolParam from typing_extensions import override +from aidial_assistant.application.request_data import PluginInfo from aidial_assistant.commands.base import ( Command, ExecutionCallback, @@ -16,7 +17,6 @@ ) from aidial_assistant.commands.open_api import OpenAPIChatCommand from aidial_assistant.commands.plugin_callback import PluginChainCallback -from aidial_assistant.commands.run_plugin import PluginInfo from aidial_assistant.model.model_client import ( ModelClient, ReasonLengthException, @@ -65,10 +65,15 @@ def _construct_tool(op: APIOperation) -> ChatCompletionToolParam: class RunTool(Command): def __init__( - self, model: ModelClient, plugin: PluginInfo, max_completion_tokens: int + self, + model: ModelClient, + plugin: PluginInfo, + auth: str | None, + max_completion_tokens: int, ): self.model = model self.plugin = plugin + self.auth = auth self.max_completion_tokens = max_completion_tokens @staticmethod @@ -86,9 +91,9 @@ async def execute( ) def create_command_tool(op: APIOperation) -> CommandTool: - return lambda: OpenAPIChatCommand( - op, self.plugin.auth - ), _construct_tool(op) + return lambda: OpenAPIChatCommand(op, self.auth), _construct_tool( + op + ) commands: CommandToolDict = { name: create_command_tool(op) for name, op in ops.items() diff --git a/aidial_assistant/model/model_client.py b/aidial_assistant/model/model_client.py index 83dfa8e..7295e35 100644 --- a/aidial_assistant/model/model_client.py +++ b/aidial_assistant/model/model_client.py @@ -1,12 +1,16 @@ from abc import ABC -from typing import Any, AsyncIterator, List +from itertools import islice +from typing import Any, AsyncIterator from aidial_sdk.utils.merge_chunks import merge from openai import AsyncOpenAI +from openai._types import NOT_GIVEN, NotGiven from openai.types.chat import ( ChatCompletionMessageParam, ChatCompletionMessageToolCallParam, + ChatCompletionToolParam, ) +from pydantic import BaseModel from aidial_assistant.utils.open_ai import Usage @@ -16,7 +20,7 @@ class ReasonLengthException(Exception): class ExtraResultsCallback: - def on_discarded_messages(self, discarded_messages: int): + def on_discarded_messages(self, discarded_messages: list[int]): pass def on_prompt_tokens(self, prompt_tokens: int): @@ -28,6 +32,13 @@ def on_tool_calls( pass +class ModelClientRequest(BaseModel): + messages: list[ChatCompletionMessageParam] + max_tokens: int | None = None + tools: list[ChatCompletionToolParam] | None = None + max_prompt_tokens: int | None = None + + async def _flush_stream(stream: AsyncIterator[str]): try: async for _ in stream: @@ -36,6 +47,25 @@ async def _flush_stream(stream: AsyncIterator[str]): pass +def _get_or_not_given(value: Any) -> Any | NotGiven: + return NOT_GIVEN if value is None else value + + +def _discarded_messages_count_to_indices( + messages: list[ChatCompletionMessageParam], discarded_messages: int +) -> list[int]: + return list( + islice( + ( + i + for i, message in enumerate(messages) + if message["role"] != "system" + ), + discarded_messages, + ) + ) + + class ModelClient(ABC): def __init__(self, client: AsyncOpenAI, model_args: dict[str, Any]): self.client = client @@ -46,15 +76,18 @@ def __init__(self, client: AsyncOpenAI, model_args: dict[str, Any]): async def agenerate( self, - messages: List[ChatCompletionMessageParam], + request: ModelClientRequest, extra_results_callback: ExtraResultsCallback | None = None, - **kwargs, ) -> AsyncIterator[str]: model_result = await self.client.chat.completions.create( **self.model_args, - extra_body=kwargs, + extra_body={"max_prompt_tokens": request.max_prompt_tokens} + if request.max_prompt_tokens is not None + else None, stream=True, - messages=messages, + messages=request.messages, + tools=_get_or_not_given(request.tools), + max_tokens=_get_or_not_given(request.max_tokens), ) finish_reason_length = False @@ -70,12 +103,16 @@ async def agenerate( extra_results_callback.on_prompt_tokens(prompt_tokens) if extra_results_callback: - discarded_messages: int | None = chunk_dict.get( + discarded_messages: int | list[int] | None = chunk_dict.get( "statistics", {} ).get("discarded_messages") if discarded_messages is not None: extra_results_callback.on_discarded_messages( - discarded_messages + _discarded_messages_count_to_indices( + request.messages, discarded_messages + ) + if isinstance(discarded_messages, int) + else discarded_messages ) choice = chunk.choices[0] @@ -106,6 +143,7 @@ async def agenerate( # TODO: Use a dedicated endpoint for counting tokens. # This request may throw an error if the number of tokens is too large. + # https://github.com/epam/ai-dial-assistant/issues/39 async def count_tokens( self, messages: list[ChatCompletionMessageParam] ) -> int: @@ -119,7 +157,8 @@ def on_prompt_tokens(self, prompt_tokens: int): callback = PromptTokensCallback() await _flush_stream( self.agenerate( - messages, extra_results_callback=callback, max_tokens=1 + ModelClientRequest(messages=messages, max_tokens=1), + extra_results_callback=callback, ) ) if callback.token_count is None: @@ -128,29 +167,32 @@ def on_prompt_tokens(self, prompt_tokens: int): return callback.token_count # TODO: Use a dedicated endpoint for discarded_messages. + # https://github.com/epam/ai-dial-assistant/issues/39 async def get_discarded_messages( self, messages: list[ChatCompletionMessageParam], max_prompt_tokens: int - ) -> int: + ) -> list[int]: class DiscardedMessagesCallback(ExtraResultsCallback): def __init__(self): - self.message_count: int | None = None + self.discarded_messages: list[int] | None = None - def on_discarded_messages(self, discarded_messages: int): - self.message_count = discarded_messages + def on_discarded_messages(self, discarded_messages: list[int]): + self.discarded_messages = discarded_messages callback = DiscardedMessagesCallback() await _flush_stream( self.agenerate( - messages, + ModelClientRequest( + messages=messages, + max_prompt_tokens=max_prompt_tokens, + max_tokens=1, + ), extra_results_callback=callback, - max_prompt_tokens=max_prompt_tokens, - max_tokens=1, ) ) - if callback.message_count is None: - raise Exception("No message count received.") + if callback.discarded_messages is None: + raise Exception("Discarded messages were not provided.") - return callback.message_count + return callback.discarded_messages @property def total_prompt_tokens(self) -> int: diff --git a/aidial_assistant/model/tokenize_client.py b/aidial_assistant/model/tokenize_client.py new file mode 100644 index 0000000..0ca31e5 --- /dev/null +++ b/aidial_assistant/model/tokenize_client.py @@ -0,0 +1,32 @@ +from typing import Any +from urllib.parse import urljoin + +from pydantic import BaseModel + +from aidial_assistant.model.model_client import ModelClientRequest +from aidial_assistant.utils.requests import apost + + +def _read_output(output: dict[str, Any]) -> int | str: + status = output["status"] + if status == "success": + return output["token_count"] + + if status == "error": + return output["error"] + + raise ValueError(f"Unknown status: {status}") + + +class TokenizeClientRequest(BaseModel): + inputs: list[ModelClientRequest | str] + + +class TokenizeClient: + def __init__(self, base_url: str): + self.url = urljoin(base_url, "tokenize") + + async def tokenize(self, request: TokenizeClientRequest) -> list[int | str]: + async with apost(self.url, request) as response: + data = await response.json() + return [_read_output(output) for output in data["outputs"]] diff --git a/aidial_assistant/model/truncate_propmt_client.py b/aidial_assistant/model/truncate_propmt_client.py new file mode 100644 index 0000000..918fc55 --- /dev/null +++ b/aidial_assistant/model/truncate_propmt_client.py @@ -0,0 +1,34 @@ +from typing import Any +from urllib.parse import urljoin + +from pydantic import BaseModel + +from aidial_assistant.model.model_client import ModelClientRequest +from aidial_assistant.utils.requests import apost + + +def _read_output(output: dict[str, Any]) -> list[int] | str: + status = output["status"] + if status == "success": + return output["discarded_messages"] + + if status == "error": + return output["error"] + + raise ValueError(f"Unknown status: {status}") + + +class TruncatePromptClientRequest(BaseModel): + inputs: list[ModelClientRequest] + + +class TruncatePromptClient: + def __init__(self, api_base: str): + self.url = urljoin(api_base, "truncate_prompt") + + async def truncate_prompt( + self, request: TruncatePromptClientRequest + ) -> list[list[int] | str]: + async with apost(self.url, request) as response: + data = await response.json() + return [_read_output(output) for output in data["outputs"]] diff --git a/aidial_assistant/tools_chain/tools_chain.py b/aidial_assistant/tools_chain/tools_chain.py index ca10121..854da20 100644 --- a/aidial_assistant/tools_chain/tools_chain.py +++ b/aidial_assistant/tools_chain/tools_chain.py @@ -33,6 +33,7 @@ from aidial_assistant.model.model_client import ( ExtraResultsCallback, ModelClient, + ModelClientRequest, ) from aidial_assistant.utils.exceptions import RequestParameterValidationError from aidial_assistant.utils.open_ai import tool_calls_message, tool_message @@ -131,11 +132,7 @@ def __init__( ): self.model = model self.commands = commands - self.model_extra_args = ( - {} - if max_completion_tokens is None - else {"max_tokens": max_completion_tokens} - ) + self.max_completion_tokens = max_completion_tokens async def run_chat( self, @@ -154,10 +151,12 @@ async def run_chat( await model_request_limiter.verify_limit(all_messages) async for chunk in self.model.agenerate( - all_messages, + ModelClientRequest( + messages=all_messages, + tools=tools, + max_tokens=self.max_completion_tokens, + ), tool_calls_callback, - tools=tools, - **self.model_extra_args, ): result_callback.on_result(chunk) except (BadRequestError, LimitExceededException) as e: @@ -172,7 +171,8 @@ async def run_chat( # and try again without tools. all_messages = all_messages[:-last_message_block_length] async for chunk in self.model.agenerate( - all_messages, tool_calls_callback + ModelClientRequest(messages=all_messages), + tool_calls_callback, ): result_callback.on_result(chunk) break diff --git a/aidial_assistant/utils/exceptions.py b/aidial_assistant/utils/exceptions.py index e8a45f1..8c52e00 100644 --- a/aidial_assistant/utils/exceptions.py +++ b/aidial_assistant/utils/exceptions.py @@ -17,6 +17,11 @@ def param(self) -> str: return self._param +class UnauthorizedAddonError(Exception): + def __init__(self, message: str, *args: object) -> None: + super().__init__(message, *args) + + def _to_http_exception(e: Exception) -> HTTPException: if isinstance(e, RequestParameterValidationError): return HTTPException( @@ -26,6 +31,11 @@ def _to_http_exception(e: Exception) -> HTTPException: param=e.param, ) + if isinstance(e, UnauthorizedAddonError): + return HTTPException( + message=str(e), status_code=401, type="invalid_request_error" + ) + if isinstance(e, APIError): raise HTTPException( message=e.message, diff --git a/aidial_assistant/utils/open_ai_plugin.py b/aidial_assistant/utils/open_ai_plugin.py index 822d197..4492a56 100644 --- a/aidial_assistant/utils/open_ai_plugin.py +++ b/aidial_assistant/utils/open_ai_plugin.py @@ -4,10 +4,8 @@ from aiocache import cached from aiohttp import hdrs -from fastapi import HTTPException from langchain.tools import OpenAPISpec from pydantic import BaseModel, parse_obj_as -from starlette.status import HTTP_401_UNAUTHORIZED from aidial_assistant.utils.requests import aget @@ -59,32 +57,6 @@ def default_auth(self) -> str | None: return self.headers.get(hdrs.AUTHORIZATION) -def get_plugin_auth( - auth_type: str, - authorization_type: str, - url: str, - token_source: AddonTokenSource, -) -> str | None: - if auth_type == "none": - return token_source.default_auth - - if auth_type == "service_http": - service_token = token_source.get_token(url) - if service_token is None: - raise HTTPException( - status_code=HTTP_401_UNAUTHORIZED, - detail=f"Missing token for {url}", - ) - - # Capitalizing because Wolfram, for instance, doesn't like lowercase bearer - return f"{authorization_type.capitalize()} {service_token}" - - raise HTTPException( - status_code=HTTP_401_UNAUTHORIZED, - detail=f"Unknown auth type {auth_type}", - ) - - async def get_open_ai_plugin_info(addon_url: str) -> OpenAIPluginInfo: """Takes url pointing to .well-known/ai-plugin.json file""" logger.info(f"Fetching plugin info from {addon_url}") diff --git a/aidial_assistant/utils/requests.py b/aidial_assistant/utils/requests.py index de785ac..21c836d 100644 --- a/aidial_assistant/utils/requests.py +++ b/aidial_assistant/utils/requests.py @@ -1,5 +1,5 @@ from contextlib import asynccontextmanager -from typing import AsyncIterator +from typing import Any, AsyncIterator from aiohttp import ClientResponse, ClientSession @@ -18,3 +18,11 @@ async def arequest( async def aget(url: str, headers=None) -> AsyncIterator[ClientResponse]: async with arequest("GET", url, headers) as response: yield response + + +@asynccontextmanager +async def apost( + url: str, data: Any, headers=None +) -> AsyncIterator[ClientResponse]: + async with arequest("POST", url, headers, data=data) as response: + yield response diff --git a/aidial_assistant/utils/state.py b/aidial_assistant/utils/state.py index f8c9dd5..5807a59 100644 --- a/aidial_assistant/utils/state.py +++ b/aidial_assistant/utils/state.py @@ -1,20 +1,5 @@ -import json from typing import TypedDict -from aidial_sdk.chat_completion.request import CustomContent, Message, Role - -from aidial_assistant.chain.command_result import ( - CommandInvocation, - commands_to_text, -) -from aidial_assistant.chain.history import MessageScope, ScopedMessage -from aidial_assistant.utils.exceptions import RequestParameterValidationError -from aidial_assistant.utils.open_ai import ( - assistant_message, - system_message, - user_message, -) - class Invocation(TypedDict): index: str | int @@ -24,85 +9,3 @@ class Invocation(TypedDict): class State(TypedDict, total=False): invocations: list[Invocation] - - -def _get_invocations(custom_content: CustomContent | None) -> list[Invocation]: - if custom_content is None: - return [] - - state: State | None = custom_content.state - if state is None: - return [] - - invocations: list[Invocation] | None = state.get("invocations") - if invocations is None: - return [] - - invocations.sort(key=lambda invocation: int(invocation["index"])) - return invocations - - -def _convert_old_commands(string: str) -> str: - """Converts old commands to new format. - Previously saved conversations with assistant will stop working if state is not updated. - - Old format: - {"commands": [{"command": "run-addon", "args": ["", ""]}]} - New format: - {"commands": [{"command": "", "arguments": {"query": ""}}]} - """ - commands = json.loads(string) - result: list[CommandInvocation] = [] - - for command in commands["commands"]: - command_name = command["command"] - # run-addon was previously called run-plugin - if command_name in ("run-addon", "run-plugin"): - args = command["args"] - result.append( - CommandInvocation(command=args[0], arguments={"query": args[1]}) - ) - else: - result.append(command) - - return commands_to_text(result) - - -def parse_history(history: list[Message]) -> list[ScopedMessage]: - messages: list[ScopedMessage] = [] - for message in history: - if message.role == Role.ASSISTANT: - invocations = _get_invocations(message.custom_content) - for invocation in invocations: - messages.append( - ScopedMessage( - scope=MessageScope.INTERNAL, - message=assistant_message( - _convert_old_commands(invocation["request"]) - ), - ) - ) - messages.append( - ScopedMessage( - scope=MessageScope.INTERNAL, - message=user_message(invocation["response"]), - ) - ) - - messages.append( - ScopedMessage(message=assistant_message(message.content or "")) - ) - elif message.role == Role.USER: - messages.append( - ScopedMessage(message=user_message(message.content or "")) - ) - elif message.role == Role.SYSTEM: - messages.append( - ScopedMessage(message=system_message(message.content or "")) - ) - else: - raise RequestParameterValidationError( - f"Role {message.role} is not supported.", param="messages" - ) - - return messages diff --git a/poetry.lock b/poetry.lock index d637bad..d786321 100644 --- a/poetry.lock +++ b/poetry.lock @@ -1,41 +1,45 @@ -# This file is automatically @generated by Poetry 1.7.1 and should not be changed by hand. +# This file is automatically @generated by Poetry 1.6.1 and should not be changed by hand. [[package]] name = "aidial-sdk" -version = "0.6.1" +version = "0.7.0rc" description = "Framework to create applications and model adapters for AI DIAL" optional = false python-versions = ">=3.8.1,<4.0" -files = [ - {file = "aidial_sdk-0.6.1-py3-none-any.whl", hash = "sha256:fe366cc5a53c8a99b742139befbd125e508195889b7852ac8981f11e2a970bb7"}, - {file = "aidial_sdk-0.6.1.tar.gz", hash = "sha256:b39fe03c17d7f427dcb87fb255935699424da1716dc64ee15116c763838f6917"}, -] +files = [] +develop = false [package.dependencies] -aiohttp = ">=3.8.3,<4.0.0" +aiohttp = "^3.8.3" fastapi = ">=0.51,<1.0" -opentelemetry-api = {version = "1.20.0", optional = true, markers = "extra == \"telemetry\""} -opentelemetry-distro = {version = "0.41b0", optional = true, markers = "extra == \"telemetry\""} -opentelemetry-exporter-otlp-proto-grpc = {version = "1.20.0", optional = true, markers = "extra == \"telemetry\""} -opentelemetry-exporter-prometheus = {version = "1.12.0rc1", optional = true, markers = "extra == \"telemetry\""} -opentelemetry-instrumentation = {version = "0.41b0", optional = true, markers = "extra == \"telemetry\""} -opentelemetry-instrumentation-aiohttp-client = {version = "0.41b0", optional = true, markers = "extra == \"telemetry\""} -opentelemetry-instrumentation-fastapi = {version = "0.41b0", optional = true, markers = "extra == \"telemetry\""} -opentelemetry-instrumentation-httpx = {version = "0.41b0", optional = true, markers = "extra == \"telemetry\""} -opentelemetry-instrumentation-logging = {version = "0.41b0", optional = true, markers = "extra == \"telemetry\""} -opentelemetry-instrumentation-requests = {version = "0.41b0", optional = true, markers = "extra == \"telemetry\""} -opentelemetry-instrumentation-system-metrics = {version = "0.41b0", optional = true, markers = "extra == \"telemetry\""} -opentelemetry-instrumentation-urllib = {version = "0.41b0", optional = true, markers = "extra == \"telemetry\""} -opentelemetry-sdk = {version = "1.20.0", optional = true, markers = "extra == \"telemetry\""} -prometheus-client = {version = "0.17.1", optional = true, markers = "extra == \"telemetry\""} +opentelemetry-api = {version = "1.20.0", optional = true} +opentelemetry-distro = {version = "0.41b0", optional = true} +opentelemetry-exporter-otlp-proto-grpc = {version = "1.20.0", optional = true} +opentelemetry-exporter-prometheus = {version = "1.12.0rc1", optional = true} +opentelemetry-instrumentation = {version = "0.41b0", optional = true} +opentelemetry-instrumentation-aiohttp-client = {version = "0.41b0", optional = true} +opentelemetry-instrumentation-fastapi = {version = "0.41b0", optional = true} +opentelemetry-instrumentation-httpx = {version = "0.41b0", optional = true} +opentelemetry-instrumentation-logging = {version = "0.41b0", optional = true} +opentelemetry-instrumentation-requests = {version = "0.41b0", optional = true} +opentelemetry-instrumentation-system-metrics = {version = "0.41b0", optional = true} +opentelemetry-instrumentation-urllib = {version = "0.41b0", optional = true} +opentelemetry-sdk = {version = "1.20.0", optional = true} +prometheus-client = {version = "0.17.1", optional = true} pydantic = ">=1.10,<3" -requests = ">=2.19,<3.0" +requests = "^2.19" uvicorn = ">=0.19,<1.0" -wrapt = ">=1.14,<2.0" +wrapt = "^1.14" [package.extras] telemetry = ["opentelemetry-api (==1.20.0)", "opentelemetry-distro (==0.41b0)", "opentelemetry-exporter-otlp-proto-grpc (==1.20.0)", "opentelemetry-exporter-prometheus (==1.12.0rc1)", "opentelemetry-instrumentation (==0.41b0)", "opentelemetry-instrumentation-aiohttp-client (==0.41b0)", "opentelemetry-instrumentation-fastapi (==0.41b0)", "opentelemetry-instrumentation-httpx (==0.41b0)", "opentelemetry-instrumentation-logging (==0.41b0)", "opentelemetry-instrumentation-requests (==0.41b0)", "opentelemetry-instrumentation-system-metrics (==0.41b0)", "opentelemetry-instrumentation-urllib (==0.41b0)", "opentelemetry-sdk (==1.20.0)", "prometheus-client (==0.17.1)"] +[package.source] +type = "git" +url = "https://github.com/epam/ai-dial-sdk.git" +reference = "feat/support-tokenize-and-truncate" +resolved_reference = "51174f2e5a2fcb8e00af22e37bee1eb862a6e341" + [[package]] name = "aiocache" version = "0.12.2" @@ -667,7 +671,7 @@ files = [ {file = "greenlet-3.0.0-cp311-cp311-musllinux_1_1_aarch64.whl", hash = "sha256:0b72b802496cccbd9b31acea72b6f87e7771ccfd7f7927437d592e5c92ed703c"}, {file = "greenlet-3.0.0-cp311-cp311-musllinux_1_1_x86_64.whl", hash = "sha256:527cd90ba3d8d7ae7dceb06fda619895768a46a1b4e423bdb24c1969823b8362"}, {file = "greenlet-3.0.0-cp311-cp311-win_amd64.whl", hash = "sha256:37f60b3a42d8b5499be910d1267b24355c495064f271cfe74bf28b17b099133c"}, - {file = "greenlet-3.0.0-cp312-cp312-macosx_10_9_universal2.whl", hash = "sha256:1482fba7fbed96ea7842b5a7fc11d61727e8be75a077e603e8ab49d24e234383"}, + {file = "greenlet-3.0.0-cp311-universal2-macosx_10_9_universal2.whl", hash = "sha256:c3692ecf3fe754c8c0f2c95ff19626584459eab110eaab66413b1e7425cd84e9"}, {file = "greenlet-3.0.0-cp312-cp312-macosx_13_0_arm64.whl", hash = "sha256:be557119bf467d37a8099d91fbf11b2de5eb1fd5fc5b91598407574848dc910f"}, {file = "greenlet-3.0.0-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:73b2f1922a39d5d59cc0e597987300df3396b148a9bd10b76a058a2f2772fc04"}, {file = "greenlet-3.0.0-cp312-cp312-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:d1e22c22f7826096ad503e9bb681b05b8c1f5a8138469b255eb91f26a76634f2"}, @@ -677,6 +681,7 @@ files = [ {file = "greenlet-3.0.0-cp312-cp312-musllinux_1_1_aarch64.whl", hash = "sha256:952256c2bc5b4ee8df8dfc54fc4de330970bf5d79253c863fb5e6761f00dda35"}, {file = "greenlet-3.0.0-cp312-cp312-musllinux_1_1_x86_64.whl", hash = "sha256:269d06fa0f9624455ce08ae0179430eea61085e3cf6457f05982b37fd2cefe17"}, {file = "greenlet-3.0.0-cp312-cp312-win_amd64.whl", hash = "sha256:9adbd8ecf097e34ada8efde9b6fec4dd2a903b1e98037adf72d12993a1c80b51"}, + {file = "greenlet-3.0.0-cp312-universal2-macosx_10_9_universal2.whl", hash = "sha256:553d6fb2324e7f4f0899e5ad2c427a4579ed4873f42124beba763f16032959af"}, {file = "greenlet-3.0.0-cp37-cp37m-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:c6b5ce7f40f0e2f8b88c28e6691ca6806814157ff05e794cdd161be928550f4c"}, {file = "greenlet-3.0.0-cp37-cp37m-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:ecf94aa539e97a8411b5ea52fc6ccd8371be9550c4041011a091eb8b3ca1d810"}, {file = "greenlet-3.0.0-cp37-cp37m-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:80dcd3c938cbcac986c5c92779db8e8ce51a89a849c135172c88ecbdc8c056b7"}, @@ -2065,14 +2070,6 @@ files = [ {file = "SQLAlchemy-2.0.21-cp311-cp311-musllinux_1_1_x86_64.whl", hash = "sha256:b69f1f754d92eb1cc6b50938359dead36b96a1dcf11a8670bff65fd9b21a4b09"}, {file = "SQLAlchemy-2.0.21-cp311-cp311-win32.whl", hash = "sha256:af520a730d523eab77d754f5cf44cc7dd7ad2d54907adeb3233177eeb22f271b"}, {file = "SQLAlchemy-2.0.21-cp311-cp311-win_amd64.whl", hash = "sha256:141675dae56522126986fa4ca713739d00ed3a6f08f3c2eb92c39c6dfec463ce"}, - {file = "SQLAlchemy-2.0.21-cp312-cp312-macosx_10_9_x86_64.whl", hash = "sha256:56628ca27aa17b5890391ded4e385bf0480209726f198799b7e980c6bd473bd7"}, - {file = "SQLAlchemy-2.0.21-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:db726be58837fe5ac39859e0fa40baafe54c6d54c02aba1d47d25536170b690f"}, - {file = "SQLAlchemy-2.0.21-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:e7421c1bfdbb7214313919472307be650bd45c4dc2fcb317d64d078993de045b"}, - {file = "SQLAlchemy-2.0.21-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:632784f7a6f12cfa0e84bf2a5003b07660addccf5563c132cd23b7cc1d7371a9"}, - {file = "SQLAlchemy-2.0.21-cp312-cp312-musllinux_1_1_aarch64.whl", hash = "sha256:f6f7276cf26145a888f2182a98f204541b519d9ea358a65d82095d9c9e22f917"}, - {file = "SQLAlchemy-2.0.21-cp312-cp312-musllinux_1_1_x86_64.whl", hash = "sha256:2a1f7ffac934bc0ea717fa1596f938483fb8c402233f9b26679b4f7b38d6ab6e"}, - {file = "SQLAlchemy-2.0.21-cp312-cp312-win32.whl", hash = "sha256:bfece2f7cec502ec5f759bbc09ce711445372deeac3628f6fa1c16b7fb45b682"}, - {file = "SQLAlchemy-2.0.21-cp312-cp312-win_amd64.whl", hash = "sha256:526b869a0f4f000d8d8ee3409d0becca30ae73f494cbb48801da0129601f72c6"}, {file = "SQLAlchemy-2.0.21-cp37-cp37m-macosx_10_9_x86_64.whl", hash = "sha256:7614f1eab4336df7dd6bee05bc974f2b02c38d3d0c78060c5faa4cd1ca2af3b8"}, {file = "SQLAlchemy-2.0.21-cp37-cp37m-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:d59cb9e20d79686aa473e0302e4a82882d7118744d30bb1dfb62d3c47141b3ec"}, {file = "SQLAlchemy-2.0.21-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:a95aa0672e3065d43c8aa80080cdd5cc40fe92dc873749e6c1cf23914c4b83af"}, @@ -2449,4 +2446,4 @@ testing = ["big-O", "jaraco.functools", "jaraco.itertools", "more-itertools", "p [metadata] lock-version = "2.0" python-versions = "^3.11" -content-hash = "c13a05722c03ba3443de34bafa7088f829bfec68e5e665df8ed3e22a584cabd0" +content-hash = "43616eabfbef2d660abd1cccb662bd76899a6d1fb389796f562ed64832db056b" diff --git a/pyproject.toml b/pyproject.toml index afa1d0b..c11228f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -24,7 +24,7 @@ openai = "^1.3.9" pydantic = "1.10.13" pyyaml = "^6.0.1" typing-extensions = "^4.8.0" -aidial-sdk = { version = "^0.6.1", extras = ["telemetry"] } +aidial-sdk = { git = "https://github.com/epam/ai-dial-sdk.git", branch = "feat/support-tokenize-and-truncate", extras = ["telemetry"] } aiohttp = "^3.9.0" openapi-schema-pydantic = "^1.2.4" openapi-pydantic = "^0.3.2" diff --git a/tests/unit_tests/chain/test_command_chain_best_effort.py b/tests/unit_tests/chain/test_command_chain_best_effort.py index f17fd4e..6f1161c 100644 --- a/tests/unit_tests/chain/test_command_chain_best_effort.py +++ b/tests/unit_tests/chain/test_command_chain_best_effort.py @@ -15,7 +15,7 @@ ) from aidial_assistant.chain.history import History, ScopedMessage from aidial_assistant.commands.base import Command, TextResult -from aidial_assistant.model.model_client import ModelClient +from aidial_assistant.model.model_client import ModelClient, ModelClientRequest from aidial_assistant.utils.open_ai import ( assistant_message, system_message, @@ -50,8 +50,8 @@ "user_message={{message}}, error={{error}}, dialogue={{dialogue}}" ), scoped_messages=[ - ScopedMessage(message=system_message(SYSTEM_MESSAGE)), - ScopedMessage(message=user_message(USER_MESSAGE)), + ScopedMessage(message=system_message(SYSTEM_MESSAGE), user_index=0), + ScopedMessage(message=user_message(USER_MESSAGE), user_index=1), ], ) @@ -82,16 +82,20 @@ async def test_model_doesnt_support_protocol(): ] assert model_client.agenerate.call_args_list == [ call( - [ - system_message(f"system_prefix={SYSTEM_MESSAGE}"), - user_message(f"{USER_MESSAGE}{ENFORCE_JSON_FORMAT}"), - ] + ModelClientRequest( + messages=[ + system_message(f"system_prefix={SYSTEM_MESSAGE}"), + user_message(f"{USER_MESSAGE}{ENFORCE_JSON_FORMAT}"), + ] + ) ), call( - [ - system_message(SYSTEM_MESSAGE), - user_message(USER_MESSAGE), - ] + ModelClientRequest( + messages=[ + system_message(SYSTEM_MESSAGE), + user_message(USER_MESSAGE), + ] + ) ), ] @@ -132,26 +136,34 @@ async def test_model_partially_supports_protocol(): ] assert model_client.agenerate.call_args_list == [ call( - [ - system_message(f"system_prefix={SYSTEM_MESSAGE}"), - user_message(f"{USER_MESSAGE}{ENFORCE_JSON_FORMAT}"), - ] + ModelClientRequest( + messages=[ + system_message(f"system_prefix={SYSTEM_MESSAGE}"), + user_message(f"{USER_MESSAGE}{ENFORCE_JSON_FORMAT}"), + ] + ) ), call( - [ - system_message(f"system_prefix={SYSTEM_MESSAGE}"), - user_message(USER_MESSAGE), - assistant_message(TEST_COMMAND_REQUEST), - user_message(f"{TEST_COMMAND_RESPONSE}{ENFORCE_JSON_FORMAT}"), - ] + ModelClientRequest( + messages=[ + system_message(f"system_prefix={SYSTEM_MESSAGE}"), + user_message(USER_MESSAGE), + assistant_message(TEST_COMMAND_REQUEST), + user_message( + f"{TEST_COMMAND_RESPONSE}{ENFORCE_JSON_FORMAT}" + ), + ] + ) ), call( - [ - system_message(SYSTEM_MESSAGE), - user_message( - f"user_message={USER_MESSAGE}, error={FAILED_PROTOCOL_ERROR}, dialogue={succeeded_dialogue}" - ), - ] + ModelClientRequest( + messages=[ + system_message(SYSTEM_MESSAGE), + user_message( + f"user_message={USER_MESSAGE}, error={FAILED_PROTOCOL_ERROR}, dialogue={succeeded_dialogue}" + ), + ] + ) ), ] @@ -197,26 +209,34 @@ async def test_no_tokens_for_tools(): ] assert model_client.agenerate.call_args_list == [ call( - [ - system_message(f"system_prefix={SYSTEM_MESSAGE}"), - user_message(f"{USER_MESSAGE}{ENFORCE_JSON_FORMAT}"), - ] + ModelClientRequest( + messages=[ + system_message(f"system_prefix={SYSTEM_MESSAGE}"), + user_message(f"{USER_MESSAGE}{ENFORCE_JSON_FORMAT}"), + ] + ) ), call( - [ - system_message(f"system_prefix={SYSTEM_MESSAGE}"), - user_message(USER_MESSAGE), - assistant_message(TEST_COMMAND_REQUEST), - user_message(f"{TEST_COMMAND_RESPONSE}{ENFORCE_JSON_FORMAT}"), - ] + ModelClientRequest( + messages=[ + system_message(f"system_prefix={SYSTEM_MESSAGE}"), + user_message(USER_MESSAGE), + assistant_message(TEST_COMMAND_REQUEST), + user_message( + f"{TEST_COMMAND_RESPONSE}{ENFORCE_JSON_FORMAT}" + ), + ] + ) ), call( - [ - system_message(SYSTEM_MESSAGE), - user_message( - f"user_message={USER_MESSAGE}, error={NO_TOKENS_ERROR}, dialogue=[]" - ), - ] + ModelClientRequest( + messages=[ + system_message(SYSTEM_MESSAGE), + user_message( + f"user_message={USER_MESSAGE}, error={NO_TOKENS_ERROR}, dialogue=[]" + ), + ] + ) ), ] @@ -255,18 +275,22 @@ async def test_model_request_limit_exceeded(): ] assert model_client.agenerate.call_args_list == [ call( - [ - system_message(f"system_prefix={SYSTEM_MESSAGE}"), - user_message(f"{USER_MESSAGE}{ENFORCE_JSON_FORMAT}"), - ] + ModelClientRequest( + messages=[ + system_message(f"system_prefix={SYSTEM_MESSAGE}"), + user_message(f"{USER_MESSAGE}{ENFORCE_JSON_FORMAT}"), + ] + ) ), call( - [ - system_message(SYSTEM_MESSAGE), - user_message( - f"user_message={USER_MESSAGE}, error={LIMIT_EXCEEDED_ERROR}, dialogue=[]" - ), - ] + ModelClientRequest( + messages=[ + system_message(SYSTEM_MESSAGE), + user_message( + f"user_message={USER_MESSAGE}, error={LIMIT_EXCEEDED_ERROR}, dialogue=[]" + ), + ] + ) ), ] assert model_request_limiter.verify_limit.call_args_list == [ diff --git a/tests/unit_tests/chain/test_history.py b/tests/unit_tests/chain/test_history.py index c3e6317..7d4186e 100644 --- a/tests/unit_tests/chain/test_history.py +++ b/tests/unit_tests/chain/test_history.py @@ -12,11 +12,11 @@ ) TRUNCATION_TEST_DATA = [ - (0, [0, 1, 2, 3, 4, 5, 6]), - (1, [0, 2, 3, 4, 5, 6]), - (2, [0, 2, 6]), - (3, [0, 2, 6]), - (4, [0, 2, 6]), + ([], [0, 1, 2, 3, 4, 5, 6]), + ([1], [0, 2, 3, 4, 5, 6]), + ([1, 3], [0, 2, 6]), + ([1, 3, 4], [0, 2, 6]), + ([1, 3, 4, 5], [0, 2, 6]), ] MAX_PROMPT_TOKENS = 123 @@ -24,92 +24,51 @@ @pytest.mark.asyncio @pytest.mark.parametrize( - "discarded_messages,expected_indices", TRUNCATION_TEST_DATA + "discarded_model_messages,expected_indices", TRUNCATION_TEST_DATA ) async def test_history_truncation( - discarded_messages: int, expected_indices: list[int] + discarded_model_messages, expected_indices: list[int] ): - history = History( + full_history = History( assistant_system_message_template=Template(""), best_effort_template=Template(""), scoped_messages=[ - ScopedMessage(message=system_message("a")), - ScopedMessage(message=user_message("b")), - ScopedMessage(message=system_message("c")), + ScopedMessage(message=system_message("a"), user_index=0), + ScopedMessage(message=user_message("b"), user_index=1), + ScopedMessage(message=system_message("c"), user_index=2), ScopedMessage( message=assistant_message("d"), scope=MessageScope.INTERNAL, + user_index=3, ), ScopedMessage( message=user_message(content="e"), scope=MessageScope.INTERNAL, + user_index=3, ), - ScopedMessage(message=assistant_message("f")), - ScopedMessage(message=user_message("g")), + ScopedMessage(message=assistant_message("f"), user_index=3), + ScopedMessage(message=user_message("g"), user_index=4), ], ) model_client = Mock(spec=ModelClient) - model_client.get_discarded_messages.return_value = discarded_messages - - actual = await history.truncate(MAX_PROMPT_TOKENS, model_client) + model_client.get_discarded_messages.return_value = discarded_model_messages - assert ( - actual.assistant_system_message_template - == history.assistant_system_message_template - ) - assert actual.best_effort_template == history.best_effort_template - assert actual.scoped_messages == [ - history.scoped_messages[i] for i in expected_indices - ] - - -@pytest.mark.asyncio -async def test_truncation_overflow(): - history = History( - assistant_system_message_template=Template(""), - best_effort_template=Template(""), - scoped_messages=[ - ScopedMessage(message=system_message("a")), - ScopedMessage(message=user_message("b")), - ], + truncated_history, _ = await full_history.truncate( + model_client, MAX_PROMPT_TOKENS ) - model_client = Mock(spec=ModelClient) - model_client.get_discarded_messages.return_value = 1 - - with pytest.raises(Exception) as exc_info: - await history.truncate(MAX_PROMPT_TOKENS, model_client) - assert ( - str(exc_info.value) == "No user messages left after history truncation." + full_history.assistant_system_message_template + == full_history.assistant_system_message_template ) - - -@pytest.mark.asyncio -async def test_truncation_with_incorrect_message_sequence(): - history = History( - assistant_system_message_template=Template(""), - best_effort_template=Template(""), - scoped_messages=[ - ScopedMessage( - message=user_message("a"), - scope=MessageScope.INTERNAL, - ), - ScopedMessage(message=user_message("b")), - ], - ) - - model_client = Mock(spec=ModelClient) - model_client.get_discarded_messages.return_value = 1 - - with pytest.raises(Exception) as exc_info: - await history.truncate(MAX_PROMPT_TOKENS, model_client) - assert ( - str(exc_info.value) - == "Internal messages must be followed by an assistant reply." + truncated_history.best_effort_template + == full_history.best_effort_template ) + assert truncated_history.scoped_messages == [ + full_history.scoped_messages[i] for i in expected_indices + ] def test_protocol_messages_with_system_message(): @@ -122,9 +81,11 @@ def test_protocol_messages_with_system_message(): ), best_effort_template=Template(""), scoped_messages=[ - ScopedMessage(message=system_message(system_content)), - ScopedMessage(message=user_message(user_content)), - ScopedMessage(message=assistant_message(assistant_content)), + ScopedMessage(message=system_message(system_content), user_index=0), + ScopedMessage(message=user_message(user_content), user_index=1), + ScopedMessage( + message=assistant_message(assistant_content), user_index=2 + ), ], ) diff --git a/tests/unit_tests/model/test_model_client.py b/tests/unit_tests/model/test_model_client.py index a5ed1cf..809bf70 100644 --- a/tests/unit_tests/model/test_model_client.py +++ b/tests/unit_tests/model/test_model_client.py @@ -1,13 +1,16 @@ +from typing import Any from unittest.mock import Mock, call import pytest from openai import AsyncOpenAI +from openai._types import NOT_GIVEN from openai.types.chat.chat_completion_chunk import ChoiceDeltaToolCall from pydantic import BaseModel from aidial_assistant.model.model_client import ( ExtraResultsCallback, ModelClient, + ModelClientRequest, ReasonLengthException, ) from aidial_assistant.utils.open_ai import ( @@ -34,7 +37,7 @@ class Choice(BaseModel): class Chunk(BaseModel): choices: list[Choice] - statistics: dict[str, int] | None = None + statistics: dict[str, Any] | None = None usage: Usage | None = None @@ -46,17 +49,21 @@ async def test_discarded_messages(): [ Chunk( choices=[Choice(delta=Delta(content=""))], - statistics={"discarded_messages": 2}, + statistics={"discarded_messages": [0, 1]}, ) ] ) model_client = ModelClient(openai_client, MODEL_ARGS) extra_results_callback = Mock(spec=ExtraResultsCallback) - await join_string(model_client.agenerate([], extra_results_callback)) + await join_string( + model_client.agenerate( + ModelClientRequest(messages=[]), extra_results_callback + ) + ) assert extra_results_callback.on_discarded_messages.call_args_list == [ - call(2) + call([0, 1]) ] @@ -73,7 +80,12 @@ async def test_content(): ) model_client = ModelClient(openai_client, MODEL_ARGS) - assert await join_string(model_client.agenerate([])) == "one, two, three" + assert ( + await join_string( + model_client.agenerate(ModelClientRequest(messages=[])) + ) + == "one, two, three" + ) @pytest.mark.asyncio @@ -97,7 +109,9 @@ async def test_reason_length_with_usage(): model_client = ModelClient(openai_client, MODEL_ARGS) with pytest.raises(ReasonLengthException): - async for chunk in model_client.agenerate([]): + async for chunk in model_client.agenerate( + ModelClientRequest(messages=[]) + ): assert chunk == "text" assert model_client.total_prompt_tokens == 1 @@ -118,17 +132,17 @@ async def test_api_args(): assistant_message("c"), ] - await join_string(model_client.agenerate(messages, extra="args")) + await join_string( + model_client.agenerate(ModelClientRequest(messages=messages)) + ) assert openai_client.chat.completions.create.call_args_list == [ call( - messages=[ - {"role": "system", "content": "a"}, - {"role": "user", "content": "b"}, - {"role": "assistant", "content": "c"}, - ], + messages=messages, **MODEL_ARGS, stream=True, - extra_body={"extra": "args"}, + tools=NOT_GIVEN, + max_tokens=NOT_GIVEN, + extra_body=None, ) ] diff --git a/tests/unit_tests/tools_chain/test_tools_chain_best_effort.py b/tests/unit_tests/tools_chain/test_tools_chain_best_effort.py index 594ad61..0490fed 100644 --- a/tests/unit_tests/tools_chain/test_tools_chain_best_effort.py +++ b/tests/unit_tests/tools_chain/test_tools_chain_best_effort.py @@ -7,6 +7,7 @@ ) from openai.types.chat.chat_completion_message_tool_call_param import Function +from aidial_assistant.model.model_client import ModelClientRequest from aidial_assistant.tools_chain.tools_chain import ToolsChain from aidial_assistant.utils.open_ai import ( construct_tool, @@ -44,9 +45,15 @@ async def test_model_request_limit_exceeded(): tool = construct_tool(TEST_COMMAND_NAME, "", {}, []) model = TestModelClient( tool_calls={ - TestModelClient.agenerate_key(messages, tools=[tool]): tool_calls + TestModelClient.agenerate_key( + ModelClientRequest(messages=messages, tools=[tool]) + ): tool_calls + }, + results={ + TestModelClient.agenerate_key( + ModelClientRequest(messages=messages) + ): BEST_EFFORT_RESPONSE }, - results={TestModelClient.agenerate_key(messages): BEST_EFFORT_RESPONSE}, ) messages_with_dialogue = messages + [ diff --git a/tests/unit_tests/utils/test_state.py b/tests/unit_tests/utils/test_state.py index 58569e8..7febe96 100644 --- a/tests/unit_tests/utils/test_state.py +++ b/tests/unit_tests/utils/test_state.py @@ -1,8 +1,8 @@ from aidial_sdk.chat_completion import CustomContent, Message, Role +from aidial_assistant.application.request_data import _parse_history from aidial_assistant.chain.history import MessageScope, ScopedMessage from aidial_assistant.utils.open_ai import assistant_message, user_message -from aidial_assistant.utils.state import parse_history FIRST_USER_MESSAGE = "" SECOND_USER_MESSAGE = "" @@ -42,37 +42,45 @@ def test_parse_history(): Message(role=Role.ASSISTANT, content=SECOND_ASSISTANT_MESSAGE), ] - assert parse_history(messages) == [ + assert _parse_history(messages) == [ ScopedMessage( scope=MessageScope.USER, message=user_message(FIRST_USER_MESSAGE), + user_index=0, ), ScopedMessage( scope=MessageScope.INTERNAL, message=assistant_message(FIRST_REQUEST_FIXED), + user_index=1, ), ScopedMessage( scope=MessageScope.INTERNAL, message=user_message(FIRST_RESPONSE), + user_index=1, ), ScopedMessage( scope=MessageScope.INTERNAL, message=assistant_message(SECOND_REQUEST), + user_index=1, ), ScopedMessage( scope=MessageScope.INTERNAL, message=user_message(content=SECOND_RESPONSE), + user_index=1, ), ScopedMessage( scope=MessageScope.USER, message=assistant_message(FIRST_ASSISTANT_MESSAGE), + user_index=1, ), ScopedMessage( scope=MessageScope.USER, message=user_message(SECOND_USER_MESSAGE), + user_index=2, ), ScopedMessage( scope=MessageScope.USER, message=assistant_message(SECOND_ASSISTANT_MESSAGE), + user_index=3, ), ] diff --git a/tests/utils/mocks.py b/tests/utils/mocks.py index a9a1941..bd7ce85 100644 --- a/tests/utils/mocks.py +++ b/tests/utils/mocks.py @@ -24,6 +24,7 @@ from aidial_assistant.model.model_client import ( ExtraResultsCallback, ModelClient, + ModelClientRequest, ) @@ -49,11 +50,10 @@ def __init__( @override async def agenerate( self, - messages: list[ChatCompletionMessageParam], + request: ModelClientRequest, extra_results_callback: ExtraResultsCallback | None = None, - **kwargs, ) -> AsyncIterator[str]: - args = TestModelClient.agenerate_key(messages, **kwargs) + args = TestModelClient.agenerate_key(request) if extra_results_callback and args in self.tool_calls: extra_results_callback.on_tool_calls(self.tool_calls[args]) return @@ -65,10 +65,8 @@ async def agenerate( assert False, f"Unexpected arguments: {args}" @staticmethod - def agenerate_key( - messages: list[ChatCompletionMessageParam], **kwargs - ) -> str: - return json.dumps({"messages": messages, **kwargs}) + def agenerate_key(request: ModelClientRequest) -> str: + return json.dumps(request.json()) class TestCommand(Command):