refactor(reme_ai): update app and ops for async execution and new LLM features

- Update app.py to use async service initialization
- Refactor multiple ops to use async_execute instead of execute
- Add support for stream and use_async flags in config
- Update LLM usage to use achat instead of chat
- Add new LLM models and update existing ones in config
- Improve error handling and logging in several ops
- Update dependencies and Python version requirements
This commit is contained in:
jinli.yl 2025-09-06 20:17:33 +08:00
parent 1e5250bbfb
commit 49c73b15b3
20 changed files with 135 additions and 63 deletions

View file

@ -13,7 +13,7 @@ authors = [
]
license = { file = "LICENSE" }
readme = "README.md"
requires-python = ">=3.12"
requires-python = ">=3.11"
classifiers = [
"Programming Language :: Python :: 3",
@ -24,7 +24,7 @@ classifiers = [
keywords = ["llm", "memory", "experience", "memoryscope", "ai", "mcp", "http"]
dependencies = [
"flowllm==0.1.3",
"flowllm==0.1.5",
]
[tool.setuptools.packages.find]

View file

@ -7,7 +7,7 @@ from reme_ai.config.config_parser import ConfigParser
def main():
with BaseService.get_service(*sys.argv[1:], parser=ConfigParser) as service:
service()
service(logo="ReMe")
if __name__ == "__main__":

View file

@ -1,4 +1,3 @@
# default config.yaml
backend: http
language: ""
thread_pool_max_workers: 32
@ -18,6 +17,9 @@ http:
flow:
retrieve_task_memory:
flow_content: build_query_op >> recall_vector_store_op >> rerank_memory_op >> rewrite_memory_op
stream: false
use_async: true
service_type: http+mcp
description: "Retrieves the most relevant top-k memory experiences from historical data based on the current query to enhance task-solving capabilities"
input_schema:
query:
@ -27,6 +29,9 @@ flow:
summary_task_memory:
flow_content: trajectory_preprocess_op >> (success_extraction_op|failure_extraction_op|comparative_extraction_op) >> memory_validation_op >> update_vector_store_op
stream: false
use_async: true
service_type: http+mcp
description: "Summarizes conversation trajectories or messages into structured memory representations for long-term storage"
input_schema:
trajectories:
@ -36,6 +41,9 @@ flow:
retrieve_personal_memory:
flow_content: set_query_op >> (extract_time_op | (retrieve_memory_op >> semantic_rank_op)) >> fuse_rerank_op
stream: false
use_async: true
service_type: http+mcp
description: "Retrieves the most relevant personal memories from historical data based on the query to enhance response quality"
input_schema:
query:
@ -45,6 +53,9 @@ flow:
summary_personal_memory:
flow_content: info_filter_op >> (get_observation_op | get_observation_with_time_op | load_today_memory_op) >> contra_repeat_op >> update_vector_store_op
stream: false
use_async: true
service_type: http+mcp
description: "Consolidates user observations and memories by filtering information and removing redundancies for efficient storage"
input_schema:
trajectories:
@ -54,6 +65,9 @@ flow:
retrieve_task_memory_simple:
flow_content: build_query_op >> recall_vector_store_op >> merge_memory_op
stream: false
use_async: true
service_type: http+mcp
description: "Retrieves the most relevant top-k memory experiences from historical data based on the current query with simplified processing"
input_schema:
query:
@ -63,6 +77,9 @@ flow:
summary_task_memory_simple:
flow_content: simple_summary_op >> update_vector_store_op
stream: false
use_async: true
service_type: http+mcp
description: "Summarizes conversation trajectories or messages into memories using a simplified approach"
input_schema:
trajectories:
@ -72,6 +89,9 @@ flow:
vector_store:
flow_content: vector_store_action_op
stream: false
use_async: false
service_type: http+mcp
description: "Directly operates on the vector store with various management actions"
input_schema:
action:
@ -142,11 +162,18 @@ op:
llm:
default:
backend: openai_compatible
# model_name: qwen3-30b-a3b-thinking-2507
model_name: qwen3-30b-a3b-instruct-2507
params:
temperature: 0.6
qwen3_30b_instruct:
backend: openai_compatible
model_name: qwen3-30b-a3b-instruct-2507
qwen3_30b_thinking:
backend: openai_compatible
model_name: qwen3-30b-a3b-thinking-2507
embedding_model:
default:
backend: openai_compatible

View file

@ -1,21 +1,24 @@
import asyncio
from flowllm import C
from flowllm.context.flow_context import FlowContext
from flowllm.op.agent.react_v2_op import ReactV2Op
from flowllm.op.llm.react_llm_op import ReactLLMOp
@C.register_op()
class SimpleReactOp(ReactV2Op):
class SimpleReactOp(ReactLLMOp):
...
if __name__ == "__main__":
async def main():
from reme_ai.config.config_parser import ConfigParser
C.set_default_service_config(parser=ConfigParser).init_by_service_config()
C.set_service_config(parser=ConfigParser, config_name="config=default").init_by_service_config()
context = FlowContext(query="茅台和五粮现在股价多少?")
op = SimpleReactOp()
op(context=context)
# from reme_ai.schema import Message
# result = op.llm.chat(messages=[Message(**{"role": "user", "content": "你叫什么名字?"})])
# print("!!!", result)
await op.async_call(context=context)
print(context.response.answer)
if __name__ == "__main__":
asyncio.run(main())

View file

@ -23,10 +23,9 @@ class ExtractTimeOp(BaseLLMOp):
"""
def get_language_value(self, value_dict: dict):
return value_dict.get(self.language, value_dict.get("en"))
def execute(self):
async def async_execute(self):
"""
Executes the primary logic of identifying and extracting time data from an LLM's response.
@ -59,7 +58,7 @@ class ExtractTimeOp(BaseLLMOp):
logger.info(f"Extracting time from query: {query[:100]}...")
# Invoke the LLM to generate a response
response = self.llm.chat([Message(role=Role.USER, content=full_prompt)])
response = await self.llm.achat([Message(role=Role.USER, content=full_prompt)])
# Handle empty or unsuccessful responses
if not response or not response.content:

View file

@ -45,7 +45,7 @@ class FuseRerankOp(BaseLLMOp):
memory.metadata["match_msg_flag"] = str(int(match_msg_flag))
return match_event_flag, match_msg_flag
def execute(self):
async def async_execute(self):
"""
Executes the reranking process on memories considering their scores, types, and temporal relevance.

View file

@ -13,7 +13,7 @@ class PrintMemoryOp(BaseOp):
"""
file_path: str = __file__
def execute(self):
async def async_execute(self):
"""
Executes the primary function, it involves:
1. Fetches the memories.

View file

@ -12,7 +12,7 @@ class ReadMessageOp(BaseOp):
"""
file_path: str = __file__
def execute(self):
async def async_execute(self):
"""
Executes the primary function to fetch unmemorized chat messages.
"""

View file

@ -1,5 +1,10 @@
from flowllm import C
from typing import List
from flowllm import C
from flowllm.schema.vector_node import VectorNode
from loguru import logger
from reme_ai.schema.memory import BaseMemory, vector_node_to_memory
from reme_ai.vector_store import RecallVectorStoreOp
@ -10,4 +15,28 @@ class RetrieveMemoryOp(RecallVectorStoreOp):
Processes these memories concurrently, sorts them by similarity, and logs the activity,
facilitating efficient memory retrieval operations within a given scope.
"""
file_path: str = __file__
async def async_execute(self):
recall_key: str = self.op_params.get("recall_key", "query")
top_k: int = self.op_params.get("top_k", 3)
query: str = self.context[recall_key]
assert query, "query should be not empty!"
workspace_id: str = self.context.workspace_id
nodes: List[VectorNode] = self.vector_store.search(query=query, workspace_id=workspace_id, top_k=top_k)
memory_list: List[BaseMemory] = []
memory_content_list: List[str] = []
for node in nodes:
memory: BaseMemory = vector_node_to_memory(node)
if memory.content not in memory_content_list:
memory_list.append(memory)
memory_content_list.append(memory.content)
logger.info(f"retrieve memory.size={len(memory_list)}")
threshold_score: float | None = self.op_params.get("threshold_score", None)
if threshold_score is not None:
memory_list = [mem for mem in memory_list if mem.score >= threshold_score or mem.score is None]
logger.info(f"after filter by threshold_score size={len(memory_list)}")
self.context.response.metadata["memory_list"] = memory_list

View file

@ -19,7 +19,7 @@ class SemanticRankOp(BaseLLMOp):
"""
file_path: str = __file__
def execute(self):
async def async_execute(self):
"""
Executes the primary workflow of the SemanticRankOp which includes:
- Retrieves query and memory list from context.
@ -56,7 +56,7 @@ class SemanticRankOp(BaseLLMOp):
logger.info(f"After deduplication: {len(memory_list)} memories")
# Perform semantic ranking using LLM
ranked_memories = self._semantic_rank_memories(query, memory_list)
ranked_memories = await self._semantic_rank_memories(query, memory_list)
if ranked_memories:
memory_list = ranked_memories
@ -71,7 +71,7 @@ class SemanticRankOp(BaseLLMOp):
# Save ranked memories back to context
self.context.response.metadata["memory_list"] = memory_list
def _semantic_rank_memories(self, query: str, memories: List[BaseMemory]) -> List[BaseMemory]:
async def _semantic_rank_memories(self, query: str, memories: List[BaseMemory]) -> List[BaseMemory]:
"""
Use LLM to semantically rank memories based on relevance to the query
"""
@ -93,7 +93,7 @@ Memories:
Please respond in JSON format:
{{"rankings": [{{"index": 0, "score": 0.8}}, {{"index": 1, "score": 0.6}}, ...]}}"""
response = self.llm.chat([Message(role=Role.USER, content=prompt)])
response = await self.llm.achat([Message(role=Role.USER, content=prompt)])
if not response or not response.content:
logger.warning("LLM ranking failed, using original order")

View file

@ -14,7 +14,7 @@ class SetQueryOp(BaseOp):
into the context, utilizing either provided parameters or details from the context.
"""
def execute(self):
async def async_execute(self):
"""
Executes the operation's primary function, which involves determining the query and its
timestamp, then storing these values within the context.

View file

@ -1,6 +1,8 @@
from typing import List, Tuple, Optional
from flowllm import C, BaseLLMOp
from flowllm.enumeration.role import Role
from flowllm.schema.message import Message as FlowMessage
from loguru import logger
from reme_ai.schema import Message, Trajectory
@ -12,7 +14,7 @@ from reme_ai.utils.op_utils import merge_messages_content, parse_json_experience
class ComparativeExtractionOp(BaseLLMOp):
file_path: str = __file__
def execute(self):
async def async_execute(self):
"""Extract comparative task memories by comparing different scoring trajectories"""
all_trajectories: List[Trajectory] = self.context.get("all_trajectories", [])
success_trajectories: List[Trajectory] = self.context.get("success_trajectories", [])
@ -26,7 +28,7 @@ class ComparativeExtractionOp(BaseLLMOp):
if highest_traj and lowest_traj and highest_traj.score > lowest_traj.score:
logger.info(
f"Extracting soft comparative task memories: highest ({highest_traj.score:.2f}) vs lowest ({lowest_traj.score:.2f})")
soft_task_memories = self._extract_soft_comparative_task_memory(highest_traj, lowest_traj)
soft_task_memories = await self._extract_soft_comparative_task_memory(highest_traj, lowest_traj)
comparative_task_memories.extend(soft_task_memories)
# Hard comparison: success vs failure (if similarity search is enabled)
@ -37,7 +39,7 @@ class ComparativeExtractionOp(BaseLLMOp):
logger.info(f"Found {len(similar_pairs)} similar pairs for hard comparison")
for success_steps, failure_steps, similarity_score in similar_pairs:
hard_task_memories = self._extract_hard_comparative_task_memory(success_steps, failure_steps,
hard_task_memories = await self._extract_hard_comparative_task_memory(success_steps, failure_steps,
similarity_score)
comparative_task_memories.extend(hard_task_memories)
@ -73,7 +75,7 @@ class ComparativeExtractionOp(BaseLLMOp):
"""Get trajectory score"""
return trajectory.score
def _extract_soft_comparative_task_memory(self, higher_traj: Trajectory, lower_traj: Trajectory) -> List[
async def _extract_soft_comparative_task_memory(self, higher_traj: Trajectory, lower_traj: Trajectory) -> List[
BaseMemory]:
"""Extract soft comparative task memory (high score vs low score)"""
higher_steps = self._get_trajectory_steps(higher_traj)
@ -105,9 +107,9 @@ class ComparativeExtractionOp(BaseLLMOp):
return task_memories
return self.llm.chat(messages=[Message(content=prompt)], callback_fn=parse_task_memories)
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)
def _extract_hard_comparative_task_memory(self, success_steps: List[Message],
async def _extract_hard_comparative_task_memory(self, success_steps: List[Message],
failure_steps: List[Message], similarity_score: float) -> List[
BaseMemory]:
"""Extract hard comparative task memory (success vs failure)"""
@ -134,7 +136,7 @@ class ComparativeExtractionOp(BaseLLMOp):
return task_memories
return self.llm.chat(messages=[Message(content=prompt)], callback_fn=parse_task_memories)
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)
@staticmethod
def _get_trajectory_steps(trajectory: Trajectory) -> List[Message]:

View file

@ -1,6 +1,8 @@
from typing import List
from flowllm import C, BaseLLMOp
from flowllm.enumeration.role import Role
from flowllm.schema.message import Message as FlowMessage
from loguru import logger
from reme_ai.schema import Message, Trajectory
@ -12,7 +14,7 @@ from reme_ai.utils.op_utils import merge_messages_content, parse_json_experience
class FailureExtractionOp(BaseLLMOp):
file_path: str = __file__
def execute(self):
async def async_execute(self):
"""Extract task memories from failed trajectories"""
failure_trajectories: List[Trajectory] = self.context.get("failure_trajectories", [])
@ -29,11 +31,11 @@ class FailureExtractionOp(BaseLLMOp):
if hasattr(trajectory, 'segments') and trajectory.segments:
# Process segmented step sequences
for segment in trajectory.segments:
task_memories = self._extract_failure_task_memory_from_steps(segment, trajectory)
task_memories = await self._extract_failure_task_memory_from_steps(segment, trajectory)
failure_task_memories.extend(task_memories)
else:
# Process entire trajectory
task_memories = self._extract_failure_task_memory_from_steps(trajectory.messages, trajectory)
task_memories = await self._extract_failure_task_memory_from_steps(trajectory.messages, trajectory)
failure_task_memories.extend(task_memories)
logger.info(f"Extracted {len(failure_task_memories)} failure task memories")
@ -41,7 +43,7 @@ class FailureExtractionOp(BaseLLMOp):
# Add task memories to context
self.context.failure_task_memories = failure_task_memories
def _extract_failure_task_memory_from_steps(self, steps: List[Message], trajectory: Trajectory) -> List[BaseMemory]:
async def _extract_failure_task_memory_from_steps(self, steps: List[Message], trajectory: Trajectory) -> List[BaseMemory]:
"""Extract task memory from failed step sequences"""
step_content = merge_messages_content(steps)
context = get_trajectory_context(trajectory, steps)
@ -70,4 +72,4 @@ class FailureExtractionOp(BaseLLMOp):
return task_memories
return self.llm.chat(messages=[Message(content=prompt)], callback_fn=parse_task_memories)
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)

View file

@ -10,7 +10,7 @@ from reme_ai.schema.memory import BaseMemory
class MemoryDeduplicationOp(BaseOp):
file_path: str = __file__
def execute(self):
async def async_execute(self):
"""Remove duplicate task memories"""
# Get task memories to deduplicate
task_memories: List[BaseMemory] = self.context.memory_list

View file

@ -3,6 +3,8 @@ import re
from typing import List, Dict, Any
from flowllm import C, BaseLLMOp
from flowllm.enumeration.role import Role
from flowllm.schema.message import Message as FlowMessage
from loguru import logger
from reme_ai.schema import Message
@ -13,7 +15,7 @@ from reme_ai.schema.memory import BaseMemory
class MemoryValidationOp(BaseLLMOp):
file_path: str = __file__
def execute(self):
async def async_execute(self):
"""Validate quality of extracted task memories"""
task_memories: List[BaseMemory] = []
@ -31,7 +33,7 @@ class MemoryValidationOp(BaseLLMOp):
validated_task_memories = []
for task_memory in task_memories:
validation_result = self._validate_single_task_memory(task_memory)
validation_result = await self._validate_single_task_memory(task_memory)
if validation_result and validation_result.get("is_valid", False):
task_memory.score = validation_result.get("score", 0.0)
validated_task_memories.append(task_memory)
@ -45,13 +47,13 @@ class MemoryValidationOp(BaseLLMOp):
self.context.response.answer = json.dumps([x.model_dump() for x in validated_task_memories])
self.context.response.metadata["memory_list"] = validated_task_memories
def _validate_single_task_memory(self, task_memory: BaseMemory) -> Dict[str, Any]:
async def _validate_single_task_memory(self, task_memory: BaseMemory) -> Dict[str, Any]:
"""Validate single task memory"""
validation_info = self._llm_validate_task_memory(task_memory)
validation_info = await self._llm_validate_task_memory(task_memory)
logger.info(f"Validating: {validation_info}")
return validation_info
def _llm_validate_task_memory(self, task_memory: BaseMemory) -> Dict[str, Any]:
async def _llm_validate_task_memory(self, task_memory: BaseMemory) -> Dict[str, Any]:
"""Validate task memory using LLM"""
try:
prompt = self.prompt_format(
@ -96,7 +98,7 @@ class MemoryValidationOp(BaseLLMOp):
"reason": f"Parse error: {str(e_inner)}"
}
return self.llm.chat(messages=[Message(content=prompt)], callback_fn=parse_validation)
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_validation)
except Exception as e:
logger.error(f"LLM validation failed: {e}")

View file

@ -2,6 +2,8 @@ import json
from typing import List, Dict
from flowllm import C, BaseLLMOp
from flowllm.enumeration.role import Role
from flowllm.schema.message import Message as FlowMessage
from loguru import logger
from reme_ai.schema import Message, Trajectory
@ -13,7 +15,7 @@ from reme_ai.utils.op_utils import merge_messages_content
class SimpleComparativeSummaryOp(BaseLLMOp):
file_path: str = __file__
def compare_summary_trajectory(self, trajectory_a: Trajectory, trajectory_b: Trajectory) -> List[BaseMemory]:
async def compare_summary_trajectory(self, trajectory_a: Trajectory, trajectory_b: Trajectory) -> List[BaseMemory]:
summary_prompt = self.prompt_format(prompt_name="summary_prompt",
execution_process_a=merge_messages_content(trajectory_a.messages),
execution_process_b=merge_messages_content(trajectory_b.messages),
@ -42,9 +44,9 @@ class SimpleComparativeSummaryOp(BaseLLMOp):
logger.exception(f"parse content failed!\n{content}")
raise e
return self.llm.chat(messages=[Message(content=summary_prompt)], callback_fn=parse_content)
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=summary_prompt)], callback_fn=parse_content)
def execute(self):
async def async_execute(self):
trajectories: list = self.context.get("trajectories", [])
trajectories: List[Trajectory] = [Trajectory(**x) if isinstance(x, dict) else x for x in trajectories]
@ -61,7 +63,7 @@ class SimpleComparativeSummaryOp(BaseLLMOp):
continue
if task_trajectories[0].score > task_trajectories[-1].score:
task_memories = self.compare_summary_trajectory(trajectory_a=task_trajectories[0],
task_memories = await self.compare_summary_trajectory(trajectory_a=task_trajectories[0],
trajectory_b=task_trajectories[-1])
memory_list.extend(task_memories)

View file

@ -2,6 +2,8 @@ import json
from typing import List
from flowllm import C, BaseLLMOp
from flowllm.enumeration.role import Role
from flowllm.schema.message import Message as FlowMessage
from loguru import logger
from reme_ai.schema import Message, Trajectory
@ -13,7 +15,7 @@ from reme_ai.utils.op_utils import merge_messages_content
class SimpleSummaryOp(BaseLLMOp):
file_path: str = __file__
def summary_trajectory(self, trajectory: Trajectory) -> List[BaseMemory]:
async def summary_trajectory(self, trajectory: Trajectory) -> List[BaseMemory]:
execution_process = merge_messages_content(trajectory.messages)
success_score_threshold: float = self.op_params.get("success_score_threshold", 0.9)
logger.info(f"success_score_threshold={success_score_threshold}")
@ -49,15 +51,15 @@ class SimpleSummaryOp(BaseLLMOp):
logger.exception(f"parse content failed!\n{content}")
raise e
return self.llm.chat(messages=[Message(content=summary_prompt)], callback_fn=parse_content)
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=summary_prompt)], callback_fn=parse_content)
def execute(self):
async def async_execute(self):
trajectories: list = self.context.trajectories
trajectories: List[Trajectory] = [Trajectory(**x) if isinstance(x, dict) else x for x in trajectories]
memory_list: List[BaseMemory] = []
for trajectory in trajectories:
memories = self.summary_trajectory(trajectory)
memories = await self.summary_trajectory(trajectory)
if memories:
memory_list.extend(memories)

View file

@ -1,6 +1,8 @@
from typing import List
from flowllm import C, BaseLLMOp
from flowllm.enumeration.role import Role
from flowllm.schema.message import Message as FlowMessage
from loguru import logger
from reme_ai.schema import Message, Trajectory
@ -12,7 +14,7 @@ from reme_ai.utils.op_utils import merge_messages_content, parse_json_experience
class SuccessExtractionOp(BaseLLMOp):
file_path: str = __file__
def execute(self):
async def async_execute(self):
"""Extract task memories from successful trajectories"""
success_trajectories: List[Trajectory] = self.context.success_trajectories
@ -29,11 +31,11 @@ class SuccessExtractionOp(BaseLLMOp):
if "segments" in trajectory.metadata:
# Process segmented step sequences
for segment in trajectory.metadata["segments"]:
task_memories = self._extract_success_task_memory_from_steps(segment, trajectory)
task_memories = await self._extract_success_task_memory_from_steps(segment, trajectory)
success_task_memories.extend(task_memories)
else:
# Process entire trajectory
task_memories = self._extract_success_task_memory_from_steps(trajectory.messages, trajectory)
task_memories = await self._extract_success_task_memory_from_steps(trajectory.messages, trajectory)
success_task_memories.extend(task_memories)
logger.info(f"Extracted {len(success_task_memories)} success task memories")
@ -41,7 +43,7 @@ class SuccessExtractionOp(BaseLLMOp):
# Add task memories to context
self.context.success_task_memories = success_task_memories
def _extract_success_task_memory_from_steps(self, steps: List[Message], trajectory: Trajectory) -> List[BaseMemory]:
async def _extract_success_task_memory_from_steps(self, steps: List[Message], trajectory: Trajectory) -> List[BaseMemory]:
"""Extract task memory from successful step sequences"""
step_content = merge_messages_content(steps)
context = get_trajectory_context(trajectory, steps)
@ -70,4 +72,4 @@ class SuccessExtractionOp(BaseLLMOp):
return task_memories
return self.llm.chat(messages=[Message(content=prompt)], callback_fn=parse_task_memories)
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)

View file

@ -10,7 +10,7 @@ from reme_ai.schema import Trajectory
class TrajectoryPreprocessOp(BaseOp):
file_path: str = __file__
def execute(self):
async def async_execute(self):
"""Preprocess trajectories: validate and classify"""
trajectories: list = self.context.get("trajectories", [])
trajectories: List[Trajectory] = [Trajectory(**x) if isinstance(x, dict) else x for x in trajectories]

View file

@ -3,6 +3,8 @@ import re
from typing import List
from flowllm import C, BaseLLMOp
from flowllm.enumeration.role import Role
from flowllm.schema.message import Message as FlowMessage
from loguru import logger
from reme_ai.schema import Message, Trajectory
@ -12,7 +14,7 @@ from reme_ai.schema import Message, Trajectory
class TrajectorySegmentationOp(BaseLLMOp):
file_path: str = __file__
def execute(self):
async def async_execute(self):
"""Segment trajectories into meaningful steps"""
# Get trajectories from context
all_trajectories: List[Trajectory] = self.context.get("all_trajectories", [])
@ -30,7 +32,7 @@ class TrajectorySegmentationOp(BaseLLMOp):
# Add segmentation info to trajectories
segmented_count = 0
for trajectory in target_trajectories:
segments = self._llm_segment_trajectory(trajectory)
segments = await self._llm_segment_trajectory(trajectory)
trajectory.metadata["segments"] = segments
segmented_count += 1
@ -51,7 +53,7 @@ class TrajectorySegmentationOp(BaseLLMOp):
else:
return all_trajectories
def _llm_segment_trajectory(self, trajectory: Trajectory) -> List[List[Message]]:
async def _llm_segment_trajectory(self, trajectory: Trajectory) -> List[List[Message]]:
"""Use LLM for trajectory segmentation"""
trajectory_content = self._format_trajectory_content(trajectory)
@ -80,7 +82,7 @@ class TrajectorySegmentationOp(BaseLLMOp):
return segments if segments else [trajectory.messages]
return self.llm.chat(messages=[Message(content=prompt)], callback_fn=parse_segmentation,
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_segmentation,
default_value=[trajectory.messages])
@staticmethod