You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

utils.py 4.1 kB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798
  1. from typing import List, Optional, Union
  2. from autogen_core.models import (
  3. AssistantMessage,
  4. FunctionExecutionResult,
  5. FunctionExecutionResultMessage,
  6. LLMMessage,
  7. UserMessage,
  8. )
  9. from typing_extensions import Literal
  10. from .messages import (
  11. FunctionCallMessage,
  12. Message,
  13. MultiModalMessage,
  14. TextMessage,
  15. )
  16. def convert_content_message_to_assistant_message(
  17. message: Union[TextMessage, MultiModalMessage, FunctionCallMessage],
  18. handle_unrepresentable: Literal["error", "ignore", "try_slice"] = "error",
  19. ) -> Optional[AssistantMessage]:
  20. match message:
  21. case TextMessage() | FunctionCallMessage():
  22. return AssistantMessage(content=message.content, source=message.source)
  23. case MultiModalMessage():
  24. if handle_unrepresentable == "error":
  25. raise ValueError("Cannot represent multimodal message as AssistantMessage")
  26. elif handle_unrepresentable == "ignore":
  27. return None
  28. elif handle_unrepresentable == "try_slice":
  29. return AssistantMessage(
  30. content="".join([x for x in message.content if isinstance(x, str)]),
  31. source=message.source,
  32. )
  33. def convert_content_message_to_user_message(
  34. message: Union[TextMessage, MultiModalMessage, FunctionCallMessage],
  35. handle_unrepresentable: Literal["error", "ignore", "try_slice"] = "error",
  36. ) -> Optional[UserMessage]:
  37. match message:
  38. case TextMessage() | MultiModalMessage():
  39. return UserMessage(content=message.content, source=message.source)
  40. case FunctionCallMessage():
  41. if handle_unrepresentable == "error":
  42. raise ValueError("Cannot represent multimodal message as UserMessage")
  43. elif handle_unrepresentable == "ignore":
  44. return None
  45. elif handle_unrepresentable == "try_slice":
  46. # TODO: what is a sliced function call?
  47. raise NotImplementedError("Sliced function calls not yet implemented")
  48. def convert_tool_call_response_message(
  49. message: FunctionExecutionResultMessage,
  50. handle_unrepresentable: Literal["error", "ignore", "try_slice"] = "error",
  51. ) -> Optional[FunctionExecutionResultMessage]:
  52. match message:
  53. case FunctionExecutionResultMessage():
  54. return FunctionExecutionResultMessage(
  55. content=[FunctionExecutionResult(content=x.content, call_id=x.call_id) for x in message.content]
  56. )
  57. def convert_messages_to_llm_messages(
  58. messages: List[Message],
  59. self_name: str,
  60. handle_unrepresentable: Literal["error", "ignore", "try_slice"] = "error",
  61. ) -> List[LLMMessage]:
  62. result: List[LLMMessage] = []
  63. for message in messages:
  64. match message:
  65. case (
  66. TextMessage(content=_, source=source)
  67. | MultiModalMessage(content=_, source=source)
  68. | FunctionCallMessage(content=_, source=source)
  69. ) if source == self_name:
  70. converted_message_1 = convert_content_message_to_assistant_message(message, handle_unrepresentable)
  71. if converted_message_1 is not None:
  72. result.append(converted_message_1)
  73. case (
  74. TextMessage(content=_, source=source)
  75. | MultiModalMessage(content=_, source=source)
  76. | FunctionCallMessage(content=_, source=source)
  77. ) if source != self_name:
  78. converted_message_2 = convert_content_message_to_user_message(message, handle_unrepresentable)
  79. if converted_message_2 is not None:
  80. result.append(converted_message_2)
  81. case FunctionExecutionResultMessage(content=_):
  82. converted_message_3 = convert_tool_call_response_message(message, handle_unrepresentable)
  83. if converted_message_3 is not None:
  84. result.append(converted_message_3)
  85. case _:
  86. raise AssertionError("unreachable")
  87. return result