From 195f97857a03f6af45cfc36c52fa9484a9848210 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Fri, 19 Jun 2026 02:04:28 +0800 Subject: [PATCH] rename reme4->reme --- benchmark/appworld/appworld_react_agent.py | 360 --- benchmark/appworld/prompt.py | 660 ----- benchmark/appworld/run_appworld.py | 197 -- benchmark/appworld/run_exp_statistic.py | 164 -- benchmark/bfcl/bfcl_agent.py | 726 ----- benchmark/bfcl/bfcl_utils.py | 399 --- benchmark/bfcl/default_ids.py | 206 -- benchmark/bfcl/init_task_memory_pool.py | 235 -- benchmark/bfcl/preprocess.py | 73 - benchmark/bfcl/run_bfcl.py | 151 - benchmark/bfcl/run_exp_statistic.py | 163 - benchmark/bfcl/split_into_trainval.py | 69 - benchmark/halumem/eval_reme.py | 1494 ---------- benchmark/locomo/eval_reme.py | 1107 ------- benchmark/longmemeval/compute_stats.py | 346 --- .../longmemeval/eval_longmemeval_reme.py | 1087 ------- .../eval_longmemeval_reme_retrieve.py | 921 ------ benchmark/longmemeval/eval_tools.py | 232 -- benchmark/longmemeval/llms.py | 97 - reme/__init__.py | 35 +- {reme4 => reme}/application.py | 0 {reme4 => reme}/components/__init__.py | 0 .../components/agent_wrapper/__init__.py | 0 .../agent_wrapper/as_agent_wrapper.py | 0 .../agent_wrapper/base_agent_wrapper.py | 0 .../agent_wrapper/cc_agent_wrapper.py | 0 .../components/application_context.py | 0 .../components/as_embedding/__init__.py | 0 {reme4 => reme}/components/as_llm/__init__.py | 0 {reme4 => reme}/components/base_component.py | 0 {reme4 => reme}/components/client/__init__.py | 0 .../components/client/base_client.py | 0 .../components/client/http_client.py | 0 .../components/client/mcp_client.py | 0 .../components/component_registry.py | 0 .../components/embedding_store/__init__.py | 0 .../embedding_store/base_embedding_store.py | 0 .../embedding_store/local_embedding_store.py | 0 .../components/file_catalog/__init__.py | 0 .../file_catalog/base_file_catalog.py | 0 .../file_catalog/local_file_catalog.py | 0 .../components/file_chunker/__init__.py | 0 .../file_chunker/base_file_chunker.py | 0 .../file_chunker/default_file_chunker.py | 0 .../file_chunker/markdown_file_chunker.py | 0 .../components/file_graph/__init__.py | 0 .../components/file_graph/base_file_graph.py | 0 .../components/file_graph/local_file_graph.py | 0 .../components/file_graph/neo4j_file_graph.py | 0 .../components/file_graph/nx_file_graph.py | 0 .../components/file_store/__init__.py | 0 .../components/file_store/base_file_store.py | 0 .../file_store/faiss_local_file_store.py | 0 .../components/file_store/local_file_store.py | 0 {reme4 => reme}/components/job/__init__.py | 0 .../components/job/background_job.py | 0 {reme4 => reme}/components/job/base_job.py | 0 {reme4 => reme}/components/job/cron_job.py | 0 {reme4 => reme}/components/job/stream_job.py | 0 .../components/keyword_index/__init__.py | 0 .../keyword_index/base_keyword_index.py | 0 .../components/keyword_index/bm25_index.py | 0 {reme4 => reme}/components/prompt_handler.py | 0 {reme4 => reme}/components/runtime_context.py | 0 .../components/service/__init__.py | 0 .../components/service/base_service.py | 0 .../components/service/http_service.py | 0 .../components/service/mcp_service.py | 0 .../components/tokenizer/__init__.py | 0 .../components/tokenizer/base_tokenizer.py | 0 .../components/tokenizer/jieba_tokenizer.py | 0 .../components/tokenizer/regex_tokenizer.py | 0 reme/config/__init__.py | 9 +- {reme4 => reme}/config/config_parser.py | 0 reme/config/reme_config_parser.py | 7 - {reme4 => reme}/constants.py | 0 reme/core/__init__.py | 51 - reme/core/application.py | 657 ----- reme/core/as_llm/__init__.py | 9 - reme/core/as_llm_formatter/__init__.py | 9 - .../reme_openai_chat_formatter.py | 215 -- reme/core/as_token_counter/__init__.py | 8 - .../as_token_counter/reme_token_counter.py | 123 - .../as_token_counter/rule_token_counter.py | 78 - reme/core/base_dict.py | 41 - reme/core/embedding/__init__.py | 15 - reme/core/embedding/base_embedding_model.py | 599 ---- reme/core/embedding/openai_embedding_model.py | 62 - .../embedding/openai_embedding_model_sync.py | 39 - reme/core/enumeration/__init__.py | 19 - reme/core/enumeration/chunk_enum.py | 31 - reme/core/enumeration/http_enum.py | 22 - reme/core/enumeration/json_schema_enum.py | 38 - reme/core/enumeration/memory_source.py | 11 - reme/core/enumeration/memory_type.py | 33 - reme/core/enumeration/registry_enum.py | 31 - reme/core/enumeration/role.py | 19 - reme/core/file_store/__init__.py | 34 - reme/core/file_store/base_file_store.py | 227 -- reme/core/file_store/chroma_file_store.py | 633 ---- reme/core/file_store/local_file_store.py | 474 --- reme/core/file_store/seekdb_file_store.py | 601 ---- reme/core/file_store/sqlite_file_store.py | 978 ------ reme/core/file_store/zvec_file_store.py | 573 ---- reme/core/file_watcher/__init__.py | 19 - reme/core/file_watcher/base_file_watcher.py | 243 -- reme/core/file_watcher/delta_file_watcher.py | 280 -- reme/core/file_watcher/full_file_watcher.py | 79 - reme/core/flow/__init__.py | 14 - reme/core/flow/base_flow.py | 208 -- reme/core/flow/cmd_flow.py | 18 - reme/core/flow/expression_flow.py | 37 - reme/core/llm/__init__.py | 21 - reme/core/llm/base_llm.py | 510 ---- reme/core/llm/lite_llm.py | 104 - reme/core/llm/lite_llm_sync.py | 45 - reme/core/llm/openai_llm.py | 113 - reme/core/llm/openai_llm_sync.py | 56 - reme/core/op/__init__.py | 24 - reme/core/op/base_op.py | 427 --- reme/core/op/base_ray_op.py | 124 - reme/core/op/base_react.py | 176 -- reme/core/op/base_react_stream.py | 249 -- reme/core/op/base_tool.py | 58 - reme/core/op/mcp_tool.py | 83 - reme/core/op/parallel_op.py | 33 - reme/core/op/sequential_op.py | 35 - reme/core/prompt_handler.py | 146 - reme/core/registry_factory.py | 50 - reme/core/runtime_context.py | 87 - reme/core/schema/__init__.py | 59 - reme/core/schema/as_msg_stat.py | 96 - reme/core/schema/cut_point_result.py | 21 - reme/core/schema/file_metadata.py | 15 - reme/core/schema/memory_chunk.py | 19 - reme/core/schema/memory_node.py | 243 -- reme/core/schema/memory_search_result.py | 23 - reme/core/schema/message.py | 177 -- reme/core/schema/request.py | 11 - reme/core/schema/response.py | 13 - reme/core/schema/service_config.py | 145 - reme/core/schema/stream_chunk.py | 14 - reme/core/schema/tool_call.py | 227 -- reme/core/schema/truncation_result.py | 35 - reme/core/schema/vector_node.py | 14 - reme/core/service/__init__.py | 18 - reme/core/service/base_service.py | 47 - reme/core/service/cmd_service.py | 42 - reme/core/service/http_service.py | 93 - reme/core/service/mcp_service.py | 57 - reme/core/service_context.py | 112 - reme/core/token_counter/__init__.py | 16 - reme/core/token_counter/base_token_counter.py | 55 - reme/core/token_counter/hf_token_counter.py | 80 - .../token_counter/openai_token_counter.py | 58 - reme/core/tools/__init__.py | 45 - reme/core/tools/execute_code.py | 40 - reme/core/tools/execute_shell.py | 46 - reme/core/tools/file/__init__.py | 0 reme/core/tools/file/base_file_tool.py | 42 - reme/core/tools/file/bash_tool.py | 185 -- reme/core/tools/file/edit_diff.py | 164 -- reme/core/tools/file/edit_tool.py | 135 - reme/core/tools/file/find_tool.py | 180 -- reme/core/tools/file/grep_tool.py | 273 -- reme/core/tools/file/ls_tool.py | 125 - reme/core/tools/file/read_tool.py | 227 -- reme/core/tools/file/truncate.py | 209 -- reme/core/tools/file/write_tool.py | 81 - reme/core/tools/search/__init__.py | 0 reme/core/tools/search/dashscope_search.py | 108 - reme/core/tools/search/mock_search.py | 62 - reme/core/tools/search/tavily_search.py | 115 - reme/core/tools/think_tool.py | 42 - reme/core/utils/__init__.py | 57 - reme/core/utils/agentscope_utils.py | 389 --- reme/core/utils/cache_handler.py | 198 -- reme/core/utils/case_converter.py | 28 - reme/core/utils/chunking_utils.py | 124 - reme/core/utils/common_utils.py | 185 -- reme/core/utils/env_utils.py | 61 - reme/core/utils/execute_utils.py | 142 - reme/core/utils/hf_token_counter_utils.py | 23 - reme/core/utils/horse.py | 165 -- reme/core/utils/http_client.py | 90 - reme/core/utils/llm_utils.py | 255 -- reme/core/utils/logger_utils.py | 59 - reme/core/utils/logo_utils.py | 85 - reme/core/utils/mcp_client.py | 121 - reme/core/utils/pydantic_config_parser.py | 198 -- reme/core/utils/pydantic_utils.py | 66 - reme/core/utils/pyseekdb_conn.py | 57 - reme/core/utils/singleton.py | 21 - reme/core/utils/std_logger.py | 122 - reme/core/utils/time.py | 93 - reme/core/vector_store/__init__.py | 51 - reme/core/vector_store/base_vector_store.py | 123 - reme/core/vector_store/chroma_vector_store.py | 444 --- reme/core/vector_store/es_vector_store.py | 535 ---- reme/core/vector_store/hologres_store.py | 633 ---- reme/core/vector_store/local_vector_store.py | 341 --- reme/core/vector_store/obvec_vector_store.py | 461 --- reme/core/vector_store/pgvector_store.py | 610 ---- reme/core/vector_store/qdrant_vector_store.py | 550 ---- reme/core/vector_store/seekdb_vector_store.py | 437 --- reme/core/vector_store/zvec_vector_store.py | 809 ----- {reme4 => reme}/enumeration/__init__.py | 0 {reme4 => reme}/enumeration/chunk_enum.py | 0 {reme4 => reme}/enumeration/component_enum.py | 0 .../enumeration/link_scope_enum.py | 0 reme/extension/__init__.py | 20 - reme/extension/procedural_memory/__init__.py | 47 - .../procedural_memory/dump_memory.py | 63 - .../procedural_memory/load_memory.py | 82 - .../procedural_memory/retrieve/__init__.py | 0 .../procedural_memory/retrieve/build_query.py | 58 - .../retrieve/memory_deletion.py | 54 - .../retrieve/memory_retrieval.py | 68 - .../retrieve/merge_memory.py | 45 - .../retrieve/rerank_memory.py | 186 -- .../retrieve/rewrite_memory.py | 201 -- .../retrieve/update_memory_metadata.py | 51 - .../procedural_memory/summary/__init__.py | 0 .../summary/comparative_extraction.py | 274 -- .../summary/failure_extraction.py | 95 - .../summary/memory_addition.py | 32 - .../summary/memory_deduplication.py | 183 -- .../summary/memory_validation.py | 139 - .../summary/success_extraction.py | 95 - .../summary/trajectory_preprocess.py | 59 - .../summary/trajectory_segmentation.py | 139 - reme/extension/procedural_memory/utils.py | 142 - reme/extension/simple_chat.py | 60 - reme/extension/stream_chat.py | 62 - reme/extension/test_op.py | 15 - reme/extension/translate_ts.py | 55 - reme/memory/__init__.py | 11 - reme/memory/file_based/__init__.py | 13 - reme/memory/file_based/components/__init__.py | 15 - reme/memory/file_based/components/cli.py | 326 -- .../memory/file_based/components/compactor.py | 115 - .../file_based/components/context_checker.py | 92 - .../file_based/components/summarizer.py | 97 - .../components/tool_result_compactor.py | 161 - .../file_based/reme_in_memory_memory.py | 299 -- reme/memory/file_based/tools/__init__.py | 13 - .../file_based/tools/browser_control.py | 2624 ----------------- .../file_based/tools/browser_snapshot.py | 248 -- reme/memory/file_based/tools/file_io.py | 360 --- reme/memory/file_based/tools/memory_get.py | 116 - reme/memory/file_based/tools/memory_search.py | 114 - reme/memory/file_based/tools/shell.py | 223 -- reme/memory/file_based/utils/__init__.py | 12 - .../memory/file_based/utils/as_msg_handler.py | 398 --- reme/memory/file_based/utils/file_utils.py | 217 -- reme/memory/vector_based/__init__.py | 33 - reme/memory/vector_based/base_memory_agent.py | 73 - reme/memory/vector_based/personal/__init__.py | 0 .../personal/personal_retriever.py | 148 - .../personal/personal_summarizer.py | 136 - .../vector_based/procedural/__init__.py | 0 .../procedural/procedural_retriever.py | 82 - .../procedural/procedural_summarizer.py | 84 - reme/memory/vector_based/reme_retriever.py | 94 - reme/memory/vector_based/reme_summarizer.py | 98 - .../memory/vector_based/tool_call/__init__.py | 0 .../vector_based/tool_call/tool_retriever.py | 83 - .../vector_based/tool_call/tool_summarizer.py | 84 - reme/memory/vector_tools/__init__.py | 66 - reme/memory/vector_tools/base_memory_tool.py | 128 - reme/memory/vector_tools/delegate_task.py | 89 - reme/memory/vector_tools/history/__init__.py | 0 .../vector_tools/history/add_history.py | 57 - .../vector_tools/history/read_history.py | 77 - .../vector_tools/history/read_history_v2.py | 123 - reme/memory/vector_tools/profiles/__init__.py | 1 - .../add_draft_and_read_all_profiles.py | 106 - .../vector_tools/profiles/add_profile.py | 70 - .../vector_tools/profiles/delete_profile.py | 52 - .../profiles/file_profile_backend.py | 234 -- .../vector_tools/profiles/profile_backend.py | 51 - .../vector_tools/profiles/profile_handler.py | 195 -- .../profiles/profile_vector_handler.py | 245 -- .../profiles/read_all_profiles.py | 54 - .../vector_tools/profiles/retrieve_profile.py | 102 - .../vector_tools/profiles/update_profile.py | 122 - .../profiles/update_profiles_v1.py | 175 -- .../profiles/vector_profile_backend.py | 55 - reme/memory/vector_tools/record/__init__.py | 0 .../record/add_and_retrieve_similar_memory.py | 127 - .../add_draft_and_retrieve_similar_memory.py | 127 - reme/memory/vector_tools/record/add_memory.py | 137 - .../vector_tools/record/delete_memory.py | 59 - .../vector_tools/record/memory_handler.py | 287 -- .../vector_tools/record/retrieve_memory.py | 140 - .../record/retrieve_recent_memory.py | 52 - .../vector_tools/record/update_memory.py | 141 - .../vector_tools/record/update_memory_v1.py | 212 -- .../vector_tools/record/update_memory_v2.py | 157 - reme/reme.py | 846 +----- reme/reme_cli.py | 192 -- reme/reme_light.py | 835 ------ {reme4 => reme}/schema/__init__.py | 0 {reme4 => reme}/schema/application_config.py | 0 {reme4 => reme}/schema/emb_node.py | 0 {reme4 => reme}/schema/file_chunk.py | 0 {reme4 => reme}/schema/file_front_matter.py | 0 {reme4 => reme}/schema/file_link.py | 0 {reme4 => reme}/schema/file_node.py | 0 {reme4 => reme}/schema/request.py | 0 {reme4 => reme}/schema/response.py | 0 {reme4 => reme}/schema/stream_chunk.py | 0 {reme4 => reme}/steps/__init__.py | 0 {reme4 => reme}/steps/base_step.py | 0 {reme4 => reme}/steps/channel/__init__.py | 0 .../steps/channel/channel_notify.py | 0 .../steps/channel/claim_channel.py | 0 {reme4 => reme}/steps/common/__init__.py | 0 {reme4 => reme}/steps/common/add.py | 0 {reme4 => reme}/steps/common/demo.py | 0 {reme4 => reme}/steps/common/health_check.py | 0 {reme4 => reme}/steps/common/help.py | 0 {reme4 => reme}/steps/common/llm_demo.py | 0 {reme4 => reme}/steps/common/stream_demo.py | 0 .../steps/common/stream_llm_demo.py | 0 {reme4 => reme}/steps/common/version.py | 0 {reme4 => reme}/steps/evolve/__init__.py | 0 {reme4 => reme}/steps/evolve/_evolve.py | 0 {reme4 => reme}/steps/evolve/auto_memory.py | 0 {reme4 => reme}/steps/evolve/auto_resource.py | 0 .../steps/evolve/dream/__init__.py | 0 {reme4 => reme}/steps/evolve/dream/extract.py | 0 {reme4 => reme}/steps/evolve/dream/finish.py | 0 .../steps/evolve/dream/integrate.py | 0 .../steps/evolve/dream/proactive.py | 0 {reme4 => reme}/steps/evolve/dream/schema.py | 0 {reme4 => reme}/steps/evolve/dream/topics.py | 0 {reme4 => reme}/steps/evolve/dream/utils.py | 0 {reme4 => reme}/steps/file_io/__init__.py | 0 {reme4 => reme}/steps/file_io/_daily_index.py | 0 {reme4 => reme}/steps/file_io/_file_io.py | 0 {reme4 => reme}/steps/file_io/_path.py | 0 {reme4 => reme}/steps/file_io/daily_create.py | 0 {reme4 => reme}/steps/file_io/daily_list.py | 0 .../steps/file_io/daily_reindex.py | 0 {reme4 => reme}/steps/file_io/delete.py | 0 {reme4 => reme}/steps/file_io/edit.py | 0 .../steps/file_io/frontmatter_delete.py | 0 .../steps/file_io/frontmatter_read.py | 0 .../steps/file_io/frontmatter_update.py | 0 {reme4 => reme}/steps/file_io/list.py | 0 {reme4 => reme}/steps/file_io/move.py | 0 {reme4 => reme}/steps/file_io/read.py | 0 {reme4 => reme}/steps/file_io/read_image.py | 0 {reme4 => reme}/steps/file_io/stat.py | 0 {reme4 => reme}/steps/file_io/write.py | 0 {reme4 => reme}/steps/index/__init__.py | 0 {reme4 => reme}/steps/index/_change_batch.py | 0 {reme4 => reme}/steps/index/_watch_rules.py | 0 {reme4 => reme}/steps/index/clear_store.py | 0 {reme4 => reme}/steps/index/init_changes.py | 0 {reme4 => reme}/steps/index/log_changes.py | 0 {reme4 => reme}/steps/index/node_search.py | 0 {reme4 => reme}/steps/index/search.py | 0 {reme4 => reme}/steps/index/traverse.py | 0 {reme4 => reme}/steps/index/update_changes.py | 0 {reme4 => reme}/steps/index/watch_changes.py | 0 {reme4 => reme}/steps/transfer/__init__.py | 0 {reme4 => reme}/steps/transfer/download.py | 0 {reme4 => reme}/steps/transfer/ingest.py | 0 {reme4 => reme}/steps/transfer/upload.py | 0 {reme4 => reme}/utils/__init__.py | 0 {reme4 => reme}/utils/agent_state_io.py | 0 {reme4 => reme}/utils/common_utils.py | 2 +- {reme4 => reme}/utils/env_utils.py | 0 {reme4 => reme}/utils/jsonl_zst.py | 0 {reme4 => reme}/utils/link_expansion.py | 0 {reme4 => reme}/utils/logger_utils.py | 0 {reme4 => reme}/utils/logo_utils.py | 0 {reme4 => reme}/utils/service_utils.py | 0 {reme4 => reme}/utils/similarity_utils.py | 0 {reme4 => reme}/utils/token_utils.py | 0 {reme4 => reme}/utils/wikilink_handler.py | 0 reme4/__init__.py | 26 - reme4/config/__init__.py | 8 - reme4/reme.py | 48 - reme_ai/__init__.py | 33 - reme_ai/agent/__init__.py | 14 - reme_ai/agent/react/__init__.py | 13 - reme_ai/agent/react/agentic_retrieve_op.py | 294 -- reme_ai/agent/react/simple_react_op.py | 70 - reme_ai/agent/tools/__init__.py | 17 - reme_ai/agent/tools/llm_mock_search_op.py | 328 --- reme_ai/agent/tools/mock_search_tools.py | 182 -- reme_ai/agent/tools/use_mock_search_op.py | 184 -- reme_ai/config/__init__.py | 14 - reme_ai/config/config_parser.py | 24 - reme_ai/constants/__init__.py | 92 - reme_ai/constants/common_constants.py | 49 - reme_ai/constants/language_constants.py | 260 -- reme_ai/enumeration/__init__.py | 13 - reme_ai/enumeration/language_enum.py | 20 - reme_ai/enumeration/working_summary_mode.py | 22 - reme_ai/main.py | 235 -- reme_ai/retrieve/__init__.py | 20 - reme_ai/retrieve/personal/__init__.py | 23 - reme_ai/retrieve/personal/extract_time_op.py | 115 - reme_ai/retrieve/personal/fuse_rerank_op.py | 228 -- reme_ai/retrieve/personal/print_memory_op.py | 149 - reme_ai/retrieve/personal/read_message_op.py | 60 - .../retrieve/personal/retrieve_memory_op.py | 61 - reme_ai/retrieve/personal/semantic_rank_op.py | 211 -- reme_ai/retrieve/personal/set_query_op.py | 44 - reme_ai/retrieve/task/__init__.py | 17 - reme_ai/retrieve/task/build_query_op.py | 61 - reme_ai/retrieve/task/merge_memory_op.py | 47 - reme_ai/retrieve/task/rerank_memory_op.py | 197 -- reme_ai/retrieve/task/rewrite_memory_op.py | 205 -- reme_ai/retrieve/tool/__init__.py | 11 - .../retrieve/tool/retrieve_tool_memory_op.py | 125 - reme_ai/retrieve/working/__init__.py | 19 - .../retrieve/working/batch_write_file_op.py | 44 - reme_ai/retrieve/working/grep_op.py | 95 - reme_ai/retrieve/working/read_file_op.py | 91 - reme_ai/retrieve/working/write_file_op.py | 89 - reme_ai/schema/__init__.py | 34 - reme_ai/schema/memory.py | 597 ---- reme_ai/service/__init__.py | 15 - .../agentscope_runtime_memory_service.py | 160 - reme_ai/service/personal_memory_service.py | 216 -- reme_ai/service/task_memory_service.py | 212 -- reme_ai/summary/__init__.py | 20 - reme_ai/summary/personal/__init__.py | 30 - reme_ai/summary/personal/contra_repeat_op.py | 160 - .../summary/personal/get_observation_op.py | 163 - .../personal/get_observation_with_time_op.py | 186 -- .../personal/get_reflection_subject_op.py | 207 -- reme_ai/summary/personal/info_filter_op.py | 200 -- .../summary/personal/load_today_memory_op.py | 127 - .../summary/personal/long_contra_repeat_op.py | 232 -- reme_ai/summary/personal/update_insight_op.py | 287 -- reme_ai/summary/task/__init__.py | 28 - .../summary/task/comparative_extraction_op.py | 280 -- reme_ai/summary/task/failure_extraction_op.py | 96 - .../summary/task/memory_deduplication_op.py | 185 -- reme_ai/summary/task/memory_validation_op.py | 131 - .../task/simple_comparative_summary_op.py | 111 - reme_ai/summary/task/simple_summary_op.py | 104 - reme_ai/summary/task/success_extraction_op.py | 96 - .../summary/task/trajectory_preprocess_op.py | 63 - .../task/trajectory_segmentation_op.py | 144 - reme_ai/summary/tool/__init__.py | 13 - .../summary/tool/parse_tool_call_result_op.py | 506 ---- .../summary/tool/summary_tool_memory_op.py | 636 ---- reme_ai/summary/working/__init__.py | 23 - reme_ai/summary/working/message_compact_op.py | 141 - .../summary/working/message_compress_op.py | 306 -- reme_ai/summary/working/message_offload_op.py | 112 - reme_ai/utils/__init__.py | 30 - reme_ai/utils/datetime_handler.py | 399 --- reme_ai/utils/op_utils.py | 188 -- reme_ai/utils/tool_memory_utils.py | 242 -- reme_ai/vector_store/__init__.py | 25 - reme_ai/vector_store/delete_memory_op.py | 57 - .../vector_store/recall_vector_store_op.py | 69 - reme_ai/vector_store/update_memory_freq_op.py | 60 - .../vector_store/update_memory_utility_op.py | 68 - .../vector_store/update_vector_store_op.py | 61 - .../vector_store/vector_store_action_op.py | 102 - test/cookbook/__init__.py | 0 test/cookbook/appworld/__init__.py | 0 .../cookbook/appworld/appworld_react_agent.py | 350 --- test/cookbook/appworld/prompt.py | 659 ----- test/cookbook/appworld/run_appworld.py | 239 -- test/cookbook/appworld/run_exp_statistic.py | 160 - test/cookbook/bfcl/__init__.py | 0 test/cookbook/bfcl/bfcl_agent.py | 707 ----- test/cookbook/bfcl/bfcl_utils.py | 395 --- test/cookbook/bfcl/init_exp_pool.py | 233 -- test/cookbook/bfcl/init_task_memory_pool.py | 233 -- test/cookbook/bfcl/local_file_to_library.py | 28 - test/cookbook/bfcl/run_bfcl.py | 126 - test/cookbook/bfcl/run_exp_statistic.py | 159 - test/cookbook/bfcl/split_into_trainval.py | 31 - test/cookbook/frozenlake/__init__.py | 0 .../frozenlake/frozenlake_react_agent.py | 378 --- test/cookbook/frozenlake/map_manager.py | 128 - test/cookbook/frozenlake/run_exp_statistic.py | 373 --- test/cookbook/frozenlake/run_frozenlake.py | 310 -- test/cookbook/simple_demo/__init__.py | 0 .../cookbook/simple_demo/import_usage_demo.py | 273 -- .../simple_demo/use_personal_memory_demo.py | 112 - .../simple_demo/use_task_memory_demo.py | 291 -- .../simple_demo/use_task_memory_mcp_demo.py | 238 -- .../simple_demo/use_tool_memory_demo.py | 256 -- test/cookbook/tool_memory/__init__.py | 0 .../tool_memory/run_reme_tool_bench.py | 651 ---- .../react_agent_with_working_memory.py | 155 - .../working_memory/work_memory_demo.py | 117 - test/test/cli/__init__.py | 18 - test/test/cli/fb_cli.py | 218 -- test/test/cli/fb_compactor.py | 114 - test/test/cli/fb_context_checker.py | 161 - test/test/cli/fb_summarizer.py | 89 - test/test/reme_cli.py | 387 --- test/test/test_fs_compactor.py | 661 ----- test/test/test_fs_context_checker.py | 291 -- test/test/test_fs_file_watch_integration.py | 521 ---- test/test/test_fs_memory_get.py | 369 --- test/test/test_fs_memory_search.py | 746 ----- test/test/test_fs_summary.py | 388 --- tests/demo_memory_search.py | 400 --- .../integration}/__init__.py | 0 .../integration/_vault_fixture.py | 12 +- .../integration/test_agent_session.py | 2 +- .../integration/test_auto_dream.py | 2 +- .../integration/test_auto_memory.py | 0 .../integration/test_auto_resource.py | 2 +- .../integration/test_embedding.py | 6 +- {tests4 => tests}/integration/test_llm.py | 2 +- .../integration/test_stream_llm.py | 8 +- tests/light/test_compactor.py | 371 --- tests/light/test_context_check.py | 1353 --------- tests/light/test_format_msgs_to_str.py | 908 ------ tests/light/test_reme4_file_io_helpers.py | 43 - tests/light/test_reme_light.py | 196 -- tests/light/test_summarizer.py | 319 -- tests/light/test_tool_result_compactor.py | 176 -- tests/light/test_tools.py | 442 --- tests/light/test_truncate_text_output.py | 381 --- tests/light/test_utils.py | 456 --- tests/test_agentscope_converter.py | 383 --- tests/test_base_context.py | 73 - tests/test_base_file_watcher.py | 799 ----- tests/test_cache_handler.py | 94 - tests/test_cache_memory_usage.py | 209 -- tests/test_chunking_utils.py | 265 -- tests/test_embedding.py | 349 --- tests/test_embedding_cache.py | 427 --- tests/test_embedding_sync.py | 348 --- tests/test_execute_utils.py | 138 - tests/test_file_store.py | 1230 -------- tests/test_fs_tool.py | 445 --- tests/test_horse.py | 81 - tests/test_keyword_search_performance.py | 165 -- tests/test_llm.py | 420 --- tests/test_llm_sync.py | 421 --- tests/test_local_file_graph.py | 53 - tests/test_local_file_store_persistence.py | 97 - tests/test_logo.py | 9 - tests/test_mcp_client.py | 127 - tests/test_mcp_server.py | 125 - tests/test_memory_vector_conversion.py | 680 ----- tests/test_message.py | 195 -- tests/test_reme.py | 118 - tests/test_reme_light_watch_paths.py | 68 - tests/test_reme_memory_error_handling.py | 168 -- tests/test_timer.py | 64 - tests/test_token_counter.py | 511 ---- tests/test_tool.py | 202 -- tests/test_tool_call.py | 360 --- tests/test_transfer_steps.py | 101 - tests/test_ts_merge.py | 60 - tests/test_vector_store.py | 2070 ------------- tests/test_zvec_vector_store.py | 906 ------ {benchmark/bfcl => tests/unit}/__init__.py | 0 {tests4 => tests}/unit/test_auto_dream.py | 10 +- .../unit/test_background_steps.py | 18 +- {tests4 => tests}/unit/test_base_component.py | 10 +- .../unit/test_bm25_index_perf.py | 4 +- {tests4 => tests}/unit/test_channel_notify.py | 8 +- {tests4 => tests}/unit/test_channel_sink.py | 2 +- {tests4 => tests}/unit/test_claim_channel.py | 10 +- {tests4 => tests}/unit/test_common_steps.py | 16 +- .../unit/test_component_registry.py | 6 +- {tests4 => tests}/unit/test_config_parser.py | 2 +- {tests4 => tests}/unit/test_cron_job.py | 8 +- {tests4 => tests}/unit/test_crud_steps.py | 10 +- {tests4 => tests}/unit/test_daily_steps.py | 4 +- .../unit/test_default_file_chunker.py | 4 +- {tests4 => tests}/unit/test_file_catalog.py | 4 +- {tests4 => tests}/unit/test_file_graph.py | 6 +- .../unit/test_file_store_consistency.py | 4 +- {tests4 => tests}/unit/test_job.py | 20 +- {tests4 => tests}/unit/test_keyword_index.py | 4 +- {tests4 => tests}/unit/test_link_expansion.py | 10 +- .../unit/test_markdown_file_chunker.py | 2 +- {tests4 => tests}/unit/test_packaging.py | 6 +- {tests4 => tests}/unit/test_prompt_handler.py | 2 +- .../unit/test_read_image_steps.py | 4 +- .../unit/test_read_with_neighbors.py | 10 +- {tests4 => tests}/unit/test_reme_cli.py | 2 +- .../unit/test_runtime_context.py | 6 +- {tests4 => tests}/unit/test_search_step.py | 10 +- {tests4 => tests}/unit/test_service.py | 4 +- {tests4 => tests}/unit/test_tokenizer.py | 2 +- {tests4 => tests}/unit/test_utils.py | 12 +- {tests4 => tests}/unit/test_wikilink_utils.py | 8 +- .../unit/test_write_metadata_lock.py | 4 +- tests/vector/test_reme_vector.py | 89 - tests4/integration/__init__.py | 0 tests4/unit/__init__.py | 0 602 files changed, 194 insertions(+), 81331 deletions(-) delete mode 100644 benchmark/appworld/appworld_react_agent.py delete mode 100644 benchmark/appworld/prompt.py delete mode 100644 benchmark/appworld/run_appworld.py delete mode 100644 benchmark/appworld/run_exp_statistic.py delete mode 100644 benchmark/bfcl/bfcl_agent.py delete mode 100644 benchmark/bfcl/bfcl_utils.py delete mode 100644 benchmark/bfcl/default_ids.py delete mode 100644 benchmark/bfcl/init_task_memory_pool.py delete mode 100644 benchmark/bfcl/preprocess.py delete mode 100644 benchmark/bfcl/run_bfcl.py delete mode 100644 benchmark/bfcl/run_exp_statistic.py delete mode 100644 benchmark/bfcl/split_into_trainval.py delete mode 100644 benchmark/halumem/eval_reme.py delete mode 100644 benchmark/locomo/eval_reme.py delete mode 100644 benchmark/longmemeval/compute_stats.py delete mode 100644 benchmark/longmemeval/eval_longmemeval_reme.py delete mode 100644 benchmark/longmemeval/eval_longmemeval_reme_retrieve.py delete mode 100644 benchmark/longmemeval/eval_tools.py delete mode 100644 benchmark/longmemeval/llms.py rename {reme4 => reme}/application.py (100%) rename {reme4 => reme}/components/__init__.py (100%) rename {reme4 => reme}/components/agent_wrapper/__init__.py (100%) rename {reme4 => reme}/components/agent_wrapper/as_agent_wrapper.py (100%) rename {reme4 => reme}/components/agent_wrapper/base_agent_wrapper.py (100%) rename {reme4 => reme}/components/agent_wrapper/cc_agent_wrapper.py (100%) rename {reme4 => reme}/components/application_context.py (100%) rename {reme4 => reme}/components/as_embedding/__init__.py (100%) rename {reme4 => reme}/components/as_llm/__init__.py (100%) rename {reme4 => reme}/components/base_component.py (100%) rename {reme4 => reme}/components/client/__init__.py (100%) rename {reme4 => reme}/components/client/base_client.py (100%) rename {reme4 => reme}/components/client/http_client.py (100%) rename {reme4 => reme}/components/client/mcp_client.py (100%) rename {reme4 => reme}/components/component_registry.py (100%) rename {reme4 => reme}/components/embedding_store/__init__.py (100%) rename {reme4 => reme}/components/embedding_store/base_embedding_store.py (100%) rename {reme4 => reme}/components/embedding_store/local_embedding_store.py (100%) rename {reme4 => reme}/components/file_catalog/__init__.py (100%) rename {reme4 => reme}/components/file_catalog/base_file_catalog.py (100%) rename {reme4 => reme}/components/file_catalog/local_file_catalog.py (100%) rename {reme4 => reme}/components/file_chunker/__init__.py (100%) rename {reme4 => reme}/components/file_chunker/base_file_chunker.py (100%) rename {reme4 => reme}/components/file_chunker/default_file_chunker.py (100%) rename {reme4 => reme}/components/file_chunker/markdown_file_chunker.py (100%) rename {reme4 => reme}/components/file_graph/__init__.py (100%) rename {reme4 => reme}/components/file_graph/base_file_graph.py (100%) rename {reme4 => reme}/components/file_graph/local_file_graph.py (100%) rename {reme4 => reme}/components/file_graph/neo4j_file_graph.py (100%) rename {reme4 => reme}/components/file_graph/nx_file_graph.py (100%) rename {reme4 => reme}/components/file_store/__init__.py (100%) rename {reme4 => reme}/components/file_store/base_file_store.py (100%) rename {reme4 => reme}/components/file_store/faiss_local_file_store.py (100%) rename {reme4 => reme}/components/file_store/local_file_store.py (100%) rename {reme4 => reme}/components/job/__init__.py (100%) rename {reme4 => reme}/components/job/background_job.py (100%) rename {reme4 => reme}/components/job/base_job.py (100%) rename {reme4 => reme}/components/job/cron_job.py (100%) rename {reme4 => reme}/components/job/stream_job.py (100%) rename {reme4 => reme}/components/keyword_index/__init__.py (100%) rename {reme4 => reme}/components/keyword_index/base_keyword_index.py (100%) rename {reme4 => reme}/components/keyword_index/bm25_index.py (100%) rename {reme4 => reme}/components/prompt_handler.py (100%) rename {reme4 => reme}/components/runtime_context.py (100%) rename {reme4 => reme}/components/service/__init__.py (100%) rename {reme4 => reme}/components/service/base_service.py (100%) rename {reme4 => reme}/components/service/http_service.py (100%) rename {reme4 => reme}/components/service/mcp_service.py (100%) rename {reme4 => reme}/components/tokenizer/__init__.py (100%) rename {reme4 => reme}/components/tokenizer/base_tokenizer.py (100%) rename {reme4 => reme}/components/tokenizer/jieba_tokenizer.py (100%) rename {reme4 => reme}/components/tokenizer/regex_tokenizer.py (100%) rename {reme4 => reme}/config/config_parser.py (100%) delete mode 100644 reme/config/reme_config_parser.py rename {reme4 => reme}/constants.py (100%) delete mode 100644 reme/core/__init__.py delete mode 100644 reme/core/application.py delete mode 100644 reme/core/as_llm/__init__.py delete mode 100644 reme/core/as_llm_formatter/__init__.py delete mode 100644 reme/core/as_llm_formatter/reme_openai_chat_formatter.py delete mode 100644 reme/core/as_token_counter/__init__.py delete mode 100644 reme/core/as_token_counter/reme_token_counter.py delete mode 100644 reme/core/as_token_counter/rule_token_counter.py delete mode 100644 reme/core/base_dict.py delete mode 100644 reme/core/embedding/__init__.py delete mode 100644 reme/core/embedding/base_embedding_model.py delete mode 100644 reme/core/embedding/openai_embedding_model.py delete mode 100644 reme/core/embedding/openai_embedding_model_sync.py delete mode 100644 reme/core/enumeration/__init__.py delete mode 100644 reme/core/enumeration/chunk_enum.py delete mode 100644 reme/core/enumeration/http_enum.py delete mode 100644 reme/core/enumeration/json_schema_enum.py delete mode 100644 reme/core/enumeration/memory_source.py delete mode 100644 reme/core/enumeration/memory_type.py delete mode 100644 reme/core/enumeration/registry_enum.py delete mode 100644 reme/core/enumeration/role.py delete mode 100644 reme/core/file_store/__init__.py delete mode 100644 reme/core/file_store/base_file_store.py delete mode 100644 reme/core/file_store/chroma_file_store.py delete mode 100644 reme/core/file_store/local_file_store.py delete mode 100644 reme/core/file_store/seekdb_file_store.py delete mode 100644 reme/core/file_store/sqlite_file_store.py delete mode 100644 reme/core/file_store/zvec_file_store.py delete mode 100644 reme/core/file_watcher/__init__.py delete mode 100644 reme/core/file_watcher/base_file_watcher.py delete mode 100644 reme/core/file_watcher/delta_file_watcher.py delete mode 100644 reme/core/file_watcher/full_file_watcher.py delete mode 100644 reme/core/flow/__init__.py delete mode 100644 reme/core/flow/base_flow.py delete mode 100644 reme/core/flow/cmd_flow.py delete mode 100644 reme/core/flow/expression_flow.py delete mode 100644 reme/core/llm/__init__.py delete mode 100644 reme/core/llm/base_llm.py delete mode 100644 reme/core/llm/lite_llm.py delete mode 100644 reme/core/llm/lite_llm_sync.py delete mode 100644 reme/core/llm/openai_llm.py delete mode 100644 reme/core/llm/openai_llm_sync.py delete mode 100644 reme/core/op/__init__.py delete mode 100644 reme/core/op/base_op.py delete mode 100644 reme/core/op/base_ray_op.py delete mode 100644 reme/core/op/base_react.py delete mode 100644 reme/core/op/base_react_stream.py delete mode 100644 reme/core/op/base_tool.py delete mode 100644 reme/core/op/mcp_tool.py delete mode 100644 reme/core/op/parallel_op.py delete mode 100644 reme/core/op/sequential_op.py delete mode 100644 reme/core/prompt_handler.py delete mode 100644 reme/core/registry_factory.py delete mode 100644 reme/core/runtime_context.py delete mode 100644 reme/core/schema/__init__.py delete mode 100644 reme/core/schema/as_msg_stat.py delete mode 100644 reme/core/schema/cut_point_result.py delete mode 100644 reme/core/schema/file_metadata.py delete mode 100644 reme/core/schema/memory_chunk.py delete mode 100644 reme/core/schema/memory_node.py delete mode 100644 reme/core/schema/memory_search_result.py delete mode 100644 reme/core/schema/message.py delete mode 100644 reme/core/schema/request.py delete mode 100644 reme/core/schema/response.py delete mode 100644 reme/core/schema/service_config.py delete mode 100644 reme/core/schema/stream_chunk.py delete mode 100644 reme/core/schema/tool_call.py delete mode 100644 reme/core/schema/truncation_result.py delete mode 100644 reme/core/schema/vector_node.py delete mode 100644 reme/core/service/__init__.py delete mode 100644 reme/core/service/base_service.py delete mode 100644 reme/core/service/cmd_service.py delete mode 100644 reme/core/service/http_service.py delete mode 100644 reme/core/service/mcp_service.py delete mode 100644 reme/core/service_context.py delete mode 100644 reme/core/token_counter/__init__.py delete mode 100644 reme/core/token_counter/base_token_counter.py delete mode 100644 reme/core/token_counter/hf_token_counter.py delete mode 100644 reme/core/token_counter/openai_token_counter.py delete mode 100644 reme/core/tools/__init__.py delete mode 100644 reme/core/tools/execute_code.py delete mode 100644 reme/core/tools/execute_shell.py delete mode 100644 reme/core/tools/file/__init__.py delete mode 100644 reme/core/tools/file/base_file_tool.py delete mode 100644 reme/core/tools/file/bash_tool.py delete mode 100644 reme/core/tools/file/edit_diff.py delete mode 100644 reme/core/tools/file/edit_tool.py delete mode 100644 reme/core/tools/file/find_tool.py delete mode 100644 reme/core/tools/file/grep_tool.py delete mode 100644 reme/core/tools/file/ls_tool.py delete mode 100644 reme/core/tools/file/read_tool.py delete mode 100644 reme/core/tools/file/truncate.py delete mode 100644 reme/core/tools/file/write_tool.py delete mode 100644 reme/core/tools/search/__init__.py delete mode 100644 reme/core/tools/search/dashscope_search.py delete mode 100644 reme/core/tools/search/mock_search.py delete mode 100644 reme/core/tools/search/tavily_search.py delete mode 100644 reme/core/tools/think_tool.py delete mode 100644 reme/core/utils/__init__.py delete mode 100644 reme/core/utils/agentscope_utils.py delete mode 100644 reme/core/utils/cache_handler.py delete mode 100644 reme/core/utils/case_converter.py delete mode 100644 reme/core/utils/chunking_utils.py delete mode 100644 reme/core/utils/common_utils.py delete mode 100644 reme/core/utils/env_utils.py delete mode 100644 reme/core/utils/execute_utils.py delete mode 100644 reme/core/utils/hf_token_counter_utils.py delete mode 100644 reme/core/utils/horse.py delete mode 100644 reme/core/utils/http_client.py delete mode 100644 reme/core/utils/llm_utils.py delete mode 100644 reme/core/utils/logger_utils.py delete mode 100644 reme/core/utils/logo_utils.py delete mode 100644 reme/core/utils/mcp_client.py delete mode 100644 reme/core/utils/pydantic_config_parser.py delete mode 100644 reme/core/utils/pydantic_utils.py delete mode 100644 reme/core/utils/pyseekdb_conn.py delete mode 100644 reme/core/utils/singleton.py delete mode 100644 reme/core/utils/std_logger.py delete mode 100644 reme/core/utils/time.py delete mode 100644 reme/core/vector_store/__init__.py delete mode 100644 reme/core/vector_store/base_vector_store.py delete mode 100644 reme/core/vector_store/chroma_vector_store.py delete mode 100644 reme/core/vector_store/es_vector_store.py delete mode 100644 reme/core/vector_store/hologres_store.py delete mode 100644 reme/core/vector_store/local_vector_store.py delete mode 100644 reme/core/vector_store/obvec_vector_store.py delete mode 100644 reme/core/vector_store/pgvector_store.py delete mode 100644 reme/core/vector_store/qdrant_vector_store.py delete mode 100644 reme/core/vector_store/seekdb_vector_store.py delete mode 100644 reme/core/vector_store/zvec_vector_store.py rename {reme4 => reme}/enumeration/__init__.py (100%) rename {reme4 => reme}/enumeration/chunk_enum.py (100%) rename {reme4 => reme}/enumeration/component_enum.py (100%) rename {reme4 => reme}/enumeration/link_scope_enum.py (100%) delete mode 100644 reme/extension/__init__.py delete mode 100644 reme/extension/procedural_memory/__init__.py delete mode 100644 reme/extension/procedural_memory/dump_memory.py delete mode 100644 reme/extension/procedural_memory/load_memory.py delete mode 100644 reme/extension/procedural_memory/retrieve/__init__.py delete mode 100644 reme/extension/procedural_memory/retrieve/build_query.py delete mode 100644 reme/extension/procedural_memory/retrieve/memory_deletion.py delete mode 100644 reme/extension/procedural_memory/retrieve/memory_retrieval.py delete mode 100644 reme/extension/procedural_memory/retrieve/merge_memory.py delete mode 100644 reme/extension/procedural_memory/retrieve/rerank_memory.py delete mode 100644 reme/extension/procedural_memory/retrieve/rewrite_memory.py delete mode 100644 reme/extension/procedural_memory/retrieve/update_memory_metadata.py delete mode 100644 reme/extension/procedural_memory/summary/__init__.py delete mode 100644 reme/extension/procedural_memory/summary/comparative_extraction.py delete mode 100644 reme/extension/procedural_memory/summary/failure_extraction.py delete mode 100644 reme/extension/procedural_memory/summary/memory_addition.py delete mode 100644 reme/extension/procedural_memory/summary/memory_deduplication.py delete mode 100644 reme/extension/procedural_memory/summary/memory_validation.py delete mode 100644 reme/extension/procedural_memory/summary/success_extraction.py delete mode 100644 reme/extension/procedural_memory/summary/trajectory_preprocess.py delete mode 100644 reme/extension/procedural_memory/summary/trajectory_segmentation.py delete mode 100644 reme/extension/procedural_memory/utils.py delete mode 100644 reme/extension/simple_chat.py delete mode 100644 reme/extension/stream_chat.py delete mode 100644 reme/extension/test_op.py delete mode 100644 reme/extension/translate_ts.py delete mode 100644 reme/memory/__init__.py delete mode 100644 reme/memory/file_based/__init__.py delete mode 100644 reme/memory/file_based/components/__init__.py delete mode 100644 reme/memory/file_based/components/cli.py delete mode 100644 reme/memory/file_based/components/compactor.py delete mode 100644 reme/memory/file_based/components/context_checker.py delete mode 100644 reme/memory/file_based/components/summarizer.py delete mode 100644 reme/memory/file_based/components/tool_result_compactor.py delete mode 100644 reme/memory/file_based/reme_in_memory_memory.py delete mode 100644 reme/memory/file_based/tools/__init__.py delete mode 100644 reme/memory/file_based/tools/browser_control.py delete mode 100644 reme/memory/file_based/tools/browser_snapshot.py delete mode 100644 reme/memory/file_based/tools/file_io.py delete mode 100644 reme/memory/file_based/tools/memory_get.py delete mode 100644 reme/memory/file_based/tools/memory_search.py delete mode 100644 reme/memory/file_based/tools/shell.py delete mode 100644 reme/memory/file_based/utils/__init__.py delete mode 100644 reme/memory/file_based/utils/as_msg_handler.py delete mode 100644 reme/memory/file_based/utils/file_utils.py delete mode 100644 reme/memory/vector_based/__init__.py delete mode 100644 reme/memory/vector_based/base_memory_agent.py delete mode 100644 reme/memory/vector_based/personal/__init__.py delete mode 100644 reme/memory/vector_based/personal/personal_retriever.py delete mode 100644 reme/memory/vector_based/personal/personal_summarizer.py delete mode 100644 reme/memory/vector_based/procedural/__init__.py delete mode 100644 reme/memory/vector_based/procedural/procedural_retriever.py delete mode 100644 reme/memory/vector_based/procedural/procedural_summarizer.py delete mode 100644 reme/memory/vector_based/reme_retriever.py delete mode 100644 reme/memory/vector_based/reme_summarizer.py delete mode 100644 reme/memory/vector_based/tool_call/__init__.py delete mode 100644 reme/memory/vector_based/tool_call/tool_retriever.py delete mode 100644 reme/memory/vector_based/tool_call/tool_summarizer.py delete mode 100644 reme/memory/vector_tools/__init__.py delete mode 100644 reme/memory/vector_tools/base_memory_tool.py delete mode 100644 reme/memory/vector_tools/delegate_task.py delete mode 100644 reme/memory/vector_tools/history/__init__.py delete mode 100644 reme/memory/vector_tools/history/add_history.py delete mode 100644 reme/memory/vector_tools/history/read_history.py delete mode 100644 reme/memory/vector_tools/history/read_history_v2.py delete mode 100644 reme/memory/vector_tools/profiles/__init__.py delete mode 100644 reme/memory/vector_tools/profiles/add_draft_and_read_all_profiles.py delete mode 100644 reme/memory/vector_tools/profiles/add_profile.py delete mode 100644 reme/memory/vector_tools/profiles/delete_profile.py delete mode 100644 reme/memory/vector_tools/profiles/file_profile_backend.py delete mode 100644 reme/memory/vector_tools/profiles/profile_backend.py delete mode 100644 reme/memory/vector_tools/profiles/profile_handler.py delete mode 100644 reme/memory/vector_tools/profiles/profile_vector_handler.py delete mode 100644 reme/memory/vector_tools/profiles/read_all_profiles.py delete mode 100644 reme/memory/vector_tools/profiles/retrieve_profile.py delete mode 100644 reme/memory/vector_tools/profiles/update_profile.py delete mode 100644 reme/memory/vector_tools/profiles/update_profiles_v1.py delete mode 100644 reme/memory/vector_tools/profiles/vector_profile_backend.py delete mode 100644 reme/memory/vector_tools/record/__init__.py delete mode 100644 reme/memory/vector_tools/record/add_and_retrieve_similar_memory.py delete mode 100644 reme/memory/vector_tools/record/add_draft_and_retrieve_similar_memory.py delete mode 100644 reme/memory/vector_tools/record/add_memory.py delete mode 100644 reme/memory/vector_tools/record/delete_memory.py delete mode 100644 reme/memory/vector_tools/record/memory_handler.py delete mode 100644 reme/memory/vector_tools/record/retrieve_memory.py delete mode 100644 reme/memory/vector_tools/record/retrieve_recent_memory.py delete mode 100644 reme/memory/vector_tools/record/update_memory.py delete mode 100644 reme/memory/vector_tools/record/update_memory_v1.py delete mode 100644 reme/memory/vector_tools/record/update_memory_v2.py delete mode 100644 reme/reme_cli.py delete mode 100644 reme/reme_light.py rename {reme4 => reme}/schema/__init__.py (100%) rename {reme4 => reme}/schema/application_config.py (100%) rename {reme4 => reme}/schema/emb_node.py (100%) rename {reme4 => reme}/schema/file_chunk.py (100%) rename {reme4 => reme}/schema/file_front_matter.py (100%) rename {reme4 => reme}/schema/file_link.py (100%) rename {reme4 => reme}/schema/file_node.py (100%) rename {reme4 => reme}/schema/request.py (100%) rename {reme4 => reme}/schema/response.py (100%) rename {reme4 => reme}/schema/stream_chunk.py (100%) rename {reme4 => reme}/steps/__init__.py (100%) rename {reme4 => reme}/steps/base_step.py (100%) rename {reme4 => reme}/steps/channel/__init__.py (100%) rename {reme4 => reme}/steps/channel/channel_notify.py (100%) rename {reme4 => reme}/steps/channel/claim_channel.py (100%) rename {reme4 => reme}/steps/common/__init__.py (100%) rename {reme4 => reme}/steps/common/add.py (100%) rename {reme4 => reme}/steps/common/demo.py (100%) rename {reme4 => reme}/steps/common/health_check.py (100%) rename {reme4 => reme}/steps/common/help.py (100%) rename {reme4 => reme}/steps/common/llm_demo.py (100%) rename {reme4 => reme}/steps/common/stream_demo.py (100%) rename {reme4 => reme}/steps/common/stream_llm_demo.py (100%) rename {reme4 => reme}/steps/common/version.py (100%) rename {reme4 => reme}/steps/evolve/__init__.py (100%) rename {reme4 => reme}/steps/evolve/_evolve.py (100%) rename {reme4 => reme}/steps/evolve/auto_memory.py (100%) rename {reme4 => reme}/steps/evolve/auto_resource.py (100%) rename {reme4 => reme}/steps/evolve/dream/__init__.py (100%) rename {reme4 => reme}/steps/evolve/dream/extract.py (100%) rename {reme4 => reme}/steps/evolve/dream/finish.py (100%) rename {reme4 => reme}/steps/evolve/dream/integrate.py (100%) rename {reme4 => reme}/steps/evolve/dream/proactive.py (100%) rename {reme4 => reme}/steps/evolve/dream/schema.py (100%) rename {reme4 => reme}/steps/evolve/dream/topics.py (100%) rename {reme4 => reme}/steps/evolve/dream/utils.py (100%) rename {reme4 => reme}/steps/file_io/__init__.py (100%) rename {reme4 => reme}/steps/file_io/_daily_index.py (100%) rename {reme4 => reme}/steps/file_io/_file_io.py (100%) rename {reme4 => reme}/steps/file_io/_path.py (100%) rename {reme4 => reme}/steps/file_io/daily_create.py (100%) rename {reme4 => reme}/steps/file_io/daily_list.py (100%) rename {reme4 => reme}/steps/file_io/daily_reindex.py (100%) rename {reme4 => reme}/steps/file_io/delete.py (100%) rename {reme4 => reme}/steps/file_io/edit.py (100%) rename {reme4 => reme}/steps/file_io/frontmatter_delete.py (100%) rename {reme4 => reme}/steps/file_io/frontmatter_read.py (100%) rename {reme4 => reme}/steps/file_io/frontmatter_update.py (100%) rename {reme4 => reme}/steps/file_io/list.py (100%) rename {reme4 => reme}/steps/file_io/move.py (100%) rename {reme4 => reme}/steps/file_io/read.py (100%) rename {reme4 => reme}/steps/file_io/read_image.py (100%) rename {reme4 => reme}/steps/file_io/stat.py (100%) rename {reme4 => reme}/steps/file_io/write.py (100%) rename {reme4 => reme}/steps/index/__init__.py (100%) rename {reme4 => reme}/steps/index/_change_batch.py (100%) rename {reme4 => reme}/steps/index/_watch_rules.py (100%) rename {reme4 => reme}/steps/index/clear_store.py (100%) rename {reme4 => reme}/steps/index/init_changes.py (100%) rename {reme4 => reme}/steps/index/log_changes.py (100%) rename {reme4 => reme}/steps/index/node_search.py (100%) rename {reme4 => reme}/steps/index/search.py (100%) rename {reme4 => reme}/steps/index/traverse.py (100%) rename {reme4 => reme}/steps/index/update_changes.py (100%) rename {reme4 => reme}/steps/index/watch_changes.py (100%) rename {reme4 => reme}/steps/transfer/__init__.py (100%) rename {reme4 => reme}/steps/transfer/download.py (100%) rename {reme4 => reme}/steps/transfer/ingest.py (100%) rename {reme4 => reme}/steps/transfer/upload.py (100%) rename {reme4 => reme}/utils/__init__.py (100%) rename {reme4 => reme}/utils/agent_state_io.py (100%) rename {reme4 => reme}/utils/common_utils.py (99%) rename {reme4 => reme}/utils/env_utils.py (100%) rename {reme4 => reme}/utils/jsonl_zst.py (100%) rename {reme4 => reme}/utils/link_expansion.py (100%) rename {reme4 => reme}/utils/logger_utils.py (100%) rename {reme4 => reme}/utils/logo_utils.py (100%) rename {reme4 => reme}/utils/service_utils.py (100%) rename {reme4 => reme}/utils/similarity_utils.py (100%) rename {reme4 => reme}/utils/token_utils.py (100%) rename {reme4 => reme}/utils/wikilink_handler.py (100%) delete mode 100644 reme4/__init__.py delete mode 100644 reme4/config/__init__.py delete mode 100644 reme4/reme.py delete mode 100644 reme_ai/__init__.py delete mode 100644 reme_ai/agent/__init__.py delete mode 100644 reme_ai/agent/react/__init__.py delete mode 100644 reme_ai/agent/react/agentic_retrieve_op.py delete mode 100644 reme_ai/agent/react/simple_react_op.py delete mode 100644 reme_ai/agent/tools/__init__.py delete mode 100644 reme_ai/agent/tools/llm_mock_search_op.py delete mode 100644 reme_ai/agent/tools/mock_search_tools.py delete mode 100644 reme_ai/agent/tools/use_mock_search_op.py delete mode 100644 reme_ai/config/__init__.py delete mode 100644 reme_ai/config/config_parser.py delete mode 100644 reme_ai/constants/__init__.py delete mode 100644 reme_ai/constants/common_constants.py delete mode 100644 reme_ai/constants/language_constants.py delete mode 100644 reme_ai/enumeration/__init__.py delete mode 100644 reme_ai/enumeration/language_enum.py delete mode 100644 reme_ai/enumeration/working_summary_mode.py delete mode 100644 reme_ai/main.py delete mode 100644 reme_ai/retrieve/__init__.py delete mode 100644 reme_ai/retrieve/personal/__init__.py delete mode 100644 reme_ai/retrieve/personal/extract_time_op.py delete mode 100644 reme_ai/retrieve/personal/fuse_rerank_op.py delete mode 100644 reme_ai/retrieve/personal/print_memory_op.py delete mode 100644 reme_ai/retrieve/personal/read_message_op.py delete mode 100644 reme_ai/retrieve/personal/retrieve_memory_op.py delete mode 100644 reme_ai/retrieve/personal/semantic_rank_op.py delete mode 100644 reme_ai/retrieve/personal/set_query_op.py delete mode 100644 reme_ai/retrieve/task/__init__.py delete mode 100644 reme_ai/retrieve/task/build_query_op.py delete mode 100644 reme_ai/retrieve/task/merge_memory_op.py delete mode 100644 reme_ai/retrieve/task/rerank_memory_op.py delete mode 100644 reme_ai/retrieve/task/rewrite_memory_op.py delete mode 100644 reme_ai/retrieve/tool/__init__.py delete mode 100644 reme_ai/retrieve/tool/retrieve_tool_memory_op.py delete mode 100644 reme_ai/retrieve/working/__init__.py delete mode 100644 reme_ai/retrieve/working/batch_write_file_op.py delete mode 100644 reme_ai/retrieve/working/grep_op.py delete mode 100644 reme_ai/retrieve/working/read_file_op.py delete mode 100644 reme_ai/retrieve/working/write_file_op.py delete mode 100644 reme_ai/schema/__init__.py delete mode 100644 reme_ai/schema/memory.py delete mode 100644 reme_ai/service/__init__.py delete mode 100644 reme_ai/service/agentscope_runtime_memory_service.py delete mode 100644 reme_ai/service/personal_memory_service.py delete mode 100644 reme_ai/service/task_memory_service.py delete mode 100644 reme_ai/summary/__init__.py delete mode 100644 reme_ai/summary/personal/__init__.py delete mode 100644 reme_ai/summary/personal/contra_repeat_op.py delete mode 100644 reme_ai/summary/personal/get_observation_op.py delete mode 100644 reme_ai/summary/personal/get_observation_with_time_op.py delete mode 100644 reme_ai/summary/personal/get_reflection_subject_op.py delete mode 100644 reme_ai/summary/personal/info_filter_op.py delete mode 100644 reme_ai/summary/personal/load_today_memory_op.py delete mode 100644 reme_ai/summary/personal/long_contra_repeat_op.py delete mode 100644 reme_ai/summary/personal/update_insight_op.py delete mode 100644 reme_ai/summary/task/__init__.py delete mode 100644 reme_ai/summary/task/comparative_extraction_op.py delete mode 100644 reme_ai/summary/task/failure_extraction_op.py delete mode 100644 reme_ai/summary/task/memory_deduplication_op.py delete mode 100644 reme_ai/summary/task/memory_validation_op.py delete mode 100644 reme_ai/summary/task/simple_comparative_summary_op.py delete mode 100644 reme_ai/summary/task/simple_summary_op.py delete mode 100644 reme_ai/summary/task/success_extraction_op.py delete mode 100644 reme_ai/summary/task/trajectory_preprocess_op.py delete mode 100644 reme_ai/summary/task/trajectory_segmentation_op.py delete mode 100644 reme_ai/summary/tool/__init__.py delete mode 100644 reme_ai/summary/tool/parse_tool_call_result_op.py delete mode 100644 reme_ai/summary/tool/summary_tool_memory_op.py delete mode 100644 reme_ai/summary/working/__init__.py delete mode 100644 reme_ai/summary/working/message_compact_op.py delete mode 100644 reme_ai/summary/working/message_compress_op.py delete mode 100644 reme_ai/summary/working/message_offload_op.py delete mode 100644 reme_ai/utils/__init__.py delete mode 100644 reme_ai/utils/datetime_handler.py delete mode 100644 reme_ai/utils/op_utils.py delete mode 100644 reme_ai/utils/tool_memory_utils.py delete mode 100644 reme_ai/vector_store/__init__.py delete mode 100644 reme_ai/vector_store/delete_memory_op.py delete mode 100644 reme_ai/vector_store/recall_vector_store_op.py delete mode 100644 reme_ai/vector_store/update_memory_freq_op.py delete mode 100644 reme_ai/vector_store/update_memory_utility_op.py delete mode 100644 reme_ai/vector_store/update_vector_store_op.py delete mode 100644 reme_ai/vector_store/vector_store_action_op.py delete mode 100644 test/cookbook/__init__.py delete mode 100644 test/cookbook/appworld/__init__.py delete mode 100644 test/cookbook/appworld/appworld_react_agent.py delete mode 100644 test/cookbook/appworld/prompt.py delete mode 100644 test/cookbook/appworld/run_appworld.py delete mode 100644 test/cookbook/appworld/run_exp_statistic.py delete mode 100644 test/cookbook/bfcl/__init__.py delete mode 100644 test/cookbook/bfcl/bfcl_agent.py delete mode 100644 test/cookbook/bfcl/bfcl_utils.py delete mode 100644 test/cookbook/bfcl/init_exp_pool.py delete mode 100644 test/cookbook/bfcl/init_task_memory_pool.py delete mode 100644 test/cookbook/bfcl/local_file_to_library.py delete mode 100644 test/cookbook/bfcl/run_bfcl.py delete mode 100644 test/cookbook/bfcl/run_exp_statistic.py delete mode 100644 test/cookbook/bfcl/split_into_trainval.py delete mode 100644 test/cookbook/frozenlake/__init__.py delete mode 100644 test/cookbook/frozenlake/frozenlake_react_agent.py delete mode 100644 test/cookbook/frozenlake/map_manager.py delete mode 100644 test/cookbook/frozenlake/run_exp_statistic.py delete mode 100644 test/cookbook/frozenlake/run_frozenlake.py delete mode 100644 test/cookbook/simple_demo/__init__.py delete mode 100644 test/cookbook/simple_demo/import_usage_demo.py delete mode 100644 test/cookbook/simple_demo/use_personal_memory_demo.py delete mode 100644 test/cookbook/simple_demo/use_task_memory_demo.py delete mode 100644 test/cookbook/simple_demo/use_task_memory_mcp_demo.py delete mode 100644 test/cookbook/simple_demo/use_tool_memory_demo.py delete mode 100644 test/cookbook/tool_memory/__init__.py delete mode 100644 test/cookbook/tool_memory/run_reme_tool_bench.py delete mode 100644 test/cookbook/working_memory/react_agent_with_working_memory.py delete mode 100644 test/cookbook/working_memory/work_memory_demo.py delete mode 100644 test/test/cli/__init__.py delete mode 100644 test/test/cli/fb_cli.py delete mode 100644 test/test/cli/fb_compactor.py delete mode 100644 test/test/cli/fb_context_checker.py delete mode 100644 test/test/cli/fb_summarizer.py delete mode 100644 test/test/reme_cli.py delete mode 100644 test/test/test_fs_compactor.py delete mode 100644 test/test/test_fs_context_checker.py delete mode 100644 test/test/test_fs_file_watch_integration.py delete mode 100644 test/test/test_fs_memory_get.py delete mode 100644 test/test/test_fs_memory_search.py delete mode 100644 test/test/test_fs_summary.py delete mode 100644 tests/demo_memory_search.py rename {benchmark/appworld => tests/integration}/__init__.py (100%) rename {tests4 => tests}/integration/_vault_fixture.py (98%) rename {tests4 => tests}/integration/test_agent_session.py (98%) rename {tests4 => tests}/integration/test_auto_dream.py (99%) rename {tests4 => tests}/integration/test_auto_memory.py (100%) rename {tests4 => tests}/integration/test_auto_resource.py (99%) rename {tests4 => tests}/integration/test_embedding.py (97%) rename {tests4 => tests}/integration/test_llm.py (98%) rename {tests4 => tests}/integration/test_stream_llm.py (97%) delete mode 100644 tests/light/test_compactor.py delete mode 100644 tests/light/test_context_check.py delete mode 100644 tests/light/test_format_msgs_to_str.py delete mode 100644 tests/light/test_reme4_file_io_helpers.py delete mode 100644 tests/light/test_reme_light.py delete mode 100644 tests/light/test_summarizer.py delete mode 100644 tests/light/test_tool_result_compactor.py delete mode 100644 tests/light/test_tools.py delete mode 100644 tests/light/test_truncate_text_output.py delete mode 100644 tests/light/test_utils.py delete mode 100644 tests/test_agentscope_converter.py delete mode 100644 tests/test_base_context.py delete mode 100644 tests/test_base_file_watcher.py delete mode 100644 tests/test_cache_handler.py delete mode 100644 tests/test_cache_memory_usage.py delete mode 100644 tests/test_chunking_utils.py delete mode 100644 tests/test_embedding.py delete mode 100644 tests/test_embedding_cache.py delete mode 100644 tests/test_embedding_sync.py delete mode 100644 tests/test_execute_utils.py delete mode 100644 tests/test_file_store.py delete mode 100644 tests/test_fs_tool.py delete mode 100644 tests/test_horse.py delete mode 100644 tests/test_keyword_search_performance.py delete mode 100644 tests/test_llm.py delete mode 100644 tests/test_llm_sync.py delete mode 100644 tests/test_local_file_graph.py delete mode 100644 tests/test_local_file_store_persistence.py delete mode 100644 tests/test_logo.py delete mode 100644 tests/test_mcp_client.py delete mode 100644 tests/test_mcp_server.py delete mode 100644 tests/test_memory_vector_conversion.py delete mode 100644 tests/test_message.py delete mode 100644 tests/test_reme.py delete mode 100644 tests/test_reme_light_watch_paths.py delete mode 100644 tests/test_reme_memory_error_handling.py delete mode 100644 tests/test_timer.py delete mode 100644 tests/test_token_counter.py delete mode 100644 tests/test_tool.py delete mode 100644 tests/test_tool_call.py delete mode 100644 tests/test_transfer_steps.py delete mode 100644 tests/test_ts_merge.py delete mode 100644 tests/test_vector_store.py delete mode 100644 tests/test_zvec_vector_store.py rename {benchmark/bfcl => tests/unit}/__init__.py (100%) rename {tests4 => tests}/unit/test_auto_dream.py (90%) rename {tests4 => tests}/unit/test_background_steps.py (98%) rename {tests4 => tests}/unit/test_base_component.py (96%) rename {tests4 => tests}/unit/test_bm25_index_perf.py (98%) rename {tests4 => tests}/unit/test_channel_notify.py (93%) rename {tests4 => tests}/unit/test_channel_sink.py (98%) rename {tests4 => tests}/unit/test_claim_channel.py (90%) rename {tests4 => tests}/unit/test_common_steps.py (95%) rename {tests4 => tests}/unit/test_component_registry.py (96%) rename {tests4 => tests}/unit/test_config_parser.py (97%) rename {tests4 => tests}/unit/test_cron_job.py (91%) rename {tests4 => tests}/unit/test_crud_steps.py (99%) rename {tests4 => tests}/unit/test_daily_steps.py (99%) rename {tests4 => tests}/unit/test_default_file_chunker.py (99%) rename {tests4 => tests}/unit/test_file_catalog.py (98%) rename {tests4 => tests}/unit/test_file_graph.py (98%) rename {tests4 => tests}/unit/test_file_store_consistency.py (97%) rename {tests4 => tests}/unit/test_job.py (95%) rename {tests4 => tests}/unit/test_keyword_index.py (99%) rename {tests4 => tests}/unit/test_link_expansion.py (96%) rename {tests4 => tests}/unit/test_markdown_file_chunker.py (99%) rename {tests4 => tests}/unit/test_packaging.py (62%) rename {tests4 => tests}/unit/test_prompt_handler.py (99%) rename {tests4 => tests}/unit/test_read_image_steps.py (98%) rename {tests4 => tests}/unit/test_read_with_neighbors.py (95%) rename {tests4 => tests}/unit/test_reme_cli.py (97%) rename {tests4 => tests}/unit/test_runtime_context.py (96%) rename {tests4 => tests}/unit/test_search_step.py (95%) rename {tests4 => tests}/unit/test_service.py (90%) rename {tests4 => tests}/unit/test_tokenizer.py (98%) rename {tests4 => tests}/unit/test_utils.py (81%) rename {tests4 => tests}/unit/test_wikilink_utils.py (98%) rename {tests4 => tests}/unit/test_write_metadata_lock.py (98%) delete mode 100644 tests/vector/test_reme_vector.py delete mode 100644 tests4/integration/__init__.py delete mode 100644 tests4/unit/__init__.py diff --git a/benchmark/appworld/appworld_react_agent.py b/benchmark/appworld/appworld_react_agent.py deleted file mode 100644 index 912f9c49..00000000 --- a/benchmark/appworld/appworld_react_agent.py +++ /dev/null @@ -1,360 +0,0 @@ -# flake8: noqa: E402, E501 -# pylint: disable=E0611 -"""A minimal ReAct Agent for AppWorld tasks.""" -import os -import re -import time -import json -import datetime -from typing import List, Any - - -import ray -import requests -from tqdm import tqdm -from loguru import logger -from openai import OpenAI -from jinja2 import Template -from dotenv import load_dotenv - -from prompt import NEW_PROMPT_TEMPLATE -from appworld import AppWorld, load_task_ids - -os.environ["APPWORLD_ROOT"] = "." - -load_dotenv("../../.env") - - -@ray.remote -class AppworldReactAgent: - """A minimal ReAct Agent for AppWorld tasks.""" - - def __init__( - self, - index: int, - task_ids: List[str], - experiment_name: str, - model_name: str = "qwen3-8b", - temperature: float = 0.9, - max_interactions: int = 30, - max_response_size: int = 129024, - num_trials: int = 1, - use_memory: bool = False, - memory_base_url: str = "http://0.0.0.0:8002/", - use_memory_addition: bool = False, - use_memory_deletion: bool = False, - delete_freq: int = 10, - freq_threshold: int = 5, - utility_threshold: float = 0.5, - ): - - self.index: int = index - self.task_ids: List[str] = task_ids - self.experiment_name: str = experiment_name - self.model_name: str = model_name - self.temperature: float = temperature - self.max_interactions: int = max_interactions - self.max_response_size: int = max_response_size - self.num_trials: int = num_trials - self.use_memory: bool = use_memory - self.use_memory_addition: bool = use_memory_addition if use_memory else False - self.use_memory_deletion: bool = use_memory_deletion if use_memory else False - self.delete_freq: int = delete_freq - self.freq_threshold: int = freq_threshold - self.utility_threshold: float = utility_threshold - - self.llm_client = OpenAI() - self.memory_base_url: str = memory_base_url - - self.history: List[List[List[dict]]] = [[] for _ in range(num_trials)] - self.retrieved_memory_list: List[List[List[Any]]] = [[] for _ in range(num_trials)] - - for run_id in range(num_trials): - for _ in range(len(task_ids)): - self.retrieved_memory_list[run_id].append([]) - self.history[run_id].append([]) - - def call_llm(self, messages: list) -> str: - """Call the LLM to generate a response to the messages.""" - for i in range(100): - try: - response = self.llm_client.chat.completions.create( - model=self.model_name, - messages=messages, - temperature=self.temperature, - extra_body={"enable_thinking": False}, - seed=0, - ) - - return response.choices[0].message.content - - except Exception as e: - logger.exception(f"encounter error with {e.args}") - time.sleep(1 + i * 10) - - return "call llm error" - - def prompt_messages(self, run_id, task_index, previous_memories: None, world: AppWorld): - """Prompt the messages to the LLM.""" - app_descriptions = json.dumps( - [{"name": k, "description": v} for (k, v) in world.task.app_descriptions.items()], - indent=1, - ) - dictionary = {"supervisor": world.task.supervisor, "app_descriptions": app_descriptions} - sys_prompt = Template(NEW_PROMPT_TEMPLATE.lstrip()).render(dictionary) - query = world.task.instruction - if self.use_memory: - if len(previous_memories) == 0: - response = self.get_memory(world.task.instruction) - if response and "memory_list" in response["metadata"]: - self.retrieved_memory_list[run_id][task_index] = response["metadata"]["memory_list"] - task_memory = re.sub(r"\bMemory\s*(\d+)\s*[:]", r"Experience \1:", response["answer"]) - logger.info(f"loaded task_memory: {task_memory}") - query = ( - "Task:\n" - + query - + "\n\nSome Related Experience to help you to complete the task:\n" - + task_memory - ) - else: - formatted_memories = [] - for i, memory in enumerate(previous_memories, 1): - condition = memory["when_to_use"] - memory_content = memory["content"] - memory_text = f"Experience {i}:\n When to use: {condition}\n Content: {memory_content}\n" - formatted_memories.append(memory_text) - query = ( - "Task:\n" - + query - + "\n\nSome Related Experience to help you to complete the task:\n" - + "\n".join(formatted_memories) - ) - messages = [ - {"role": "system", "content": sys_prompt}, - {"role": "user", "content": query}, - ] - self.history[run_id][task_index] = messages - - @staticmethod - def get_reward(world) -> float: - """Get the reward for the Appworld world.""" - tracker = world.evaluate() - num_passes = len(tracker.passes) - num_failures = len(tracker.failures) - return num_passes / (num_passes + num_failures) - - def extract_code_and_fix_content( - self, - text: str, - ignore_multiple_calls=True, - ) -> tuple[str, str]: - """Extract the code and fix the content.""" - full_code_regex = r"```python\n(.*?)```" - partial_code_regex = r".*```python\n(.*)" - - original_text = text - output_code = "" - match_end = 0 - # Handle multiple calls - for re_match in re.finditer(full_code_regex, original_text, flags=re.DOTALL): - code = re_match.group(1).strip() - if ignore_multiple_calls: - text = original_text[: re_match.end()] - return code, text - output_code += code + "\n" - match_end = re_match.end() - # check for partial code match at end (no terminating ```) following the last match - partial_match = re.match( - partial_code_regex, - original_text[match_end:], - flags=re.DOTALL, - ) - if partial_match: - output_code += partial_match.group(1).strip() - # terminated due to stop condition. Add stop condition to output. - if not text.endswith("\n"): - text = text + "\n" - text = text + "```" - if len(output_code) == 0: - return text, text - else: - return output_code, text - - def execute(self): - """Execute the Appworld tasks.""" - result = [] - counter = 0 - for task_index, task_id in enumerate(tqdm(self.task_ids, desc=f"run_index={self.index}")): - t_result = None - previous_memories = [] - # Run each task num_trials times - for run_id in range(self.num_trials): - start_time = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") - with AppWorld(task_id=task_id, experiment_name=f"{self.experiment_name}_run_{run_id}") as world: - before_score = self.get_reward(world) - for i in range(self.max_interactions): - if i == 0: - self.prompt_messages( - run_id=run_id, - task_index=task_index, - previous_memories=previous_memories, - world=world, - ) - code_msg = self.call_llm(self.history[run_id][task_index]) - code, _ = self.extract_code_and_fix_content(code_msg) - self.history[run_id][task_index].append({"role": "assistant", "content": code}) - - output = world.execute(code) - # if len(output) > self.max_response_size: - # # logger.warning(f"output exceed max size={len(output)}") - # output = output[: self.max_response_size] - self.history[run_id][task_index].append( - {"role": "user", "content": "Output:\n```\n" + output + "```\n\n"}, - ) - - if world.task_completed(): - break - - after_score = self.get_reward(world) - uplift_score = after_score - before_score - - if self.use_memory: - if self.use_memory_addition: - new_traj_list = [ - self.get_traj_from_task_history(task_id, self.history[run_id][task_index], after_score), - ] - previous_memories = self.summary_memory(new_traj_list) - if after_score == 1: - self.add_memory(previous_memories) - - # update the freq & utility attributes of retrieved memories - update_utility: bool = after_score == 1 - self.update_memory_information(self.retrieved_memory_list[run_id][task_index], update_utility) - - counter += 1 - if self.use_memory_deletion: # and counter % self.delete_freq == 0: - self.delete_memory() - - t_result = { - "task_id": world.task_id, - "run_id": run_id, - "experiment_name": self.experiment_name, - "task_completed": world.task_completed(), - "before_score": before_score, - "after_score": after_score, - "uplift_score": uplift_score, - "task_history": self.history[run_id][task_index], - "task_start_time": start_time, - } - if after_score == 1: - break - result.append(t_result) - - return result - - def handle_api_response(self, response: requests.Response): - """Handle API response with proper error checking""" - if response.status_code != 200: - print(f"Error: {response.status_code}") - print(response.text) - return None - - return response.json() - - def get_memory(self, query: str): - """Retrieve relevant task memories based on a query""" - response = requests.post( - url=f"{self.memory_base_url}retrieve_task_memory", - json={ - "query": query, - "enable_llm_rerank": False, - "enable_score_filter": False, - "top_k": 5, - "enable_llm_rewrite": False, - }, - ) - - result = self.handle_api_response(response) - if not result: - return None - - logger.info(f"query: {query}, response: {result}") - return result - - def get_traj_from_task_history(self, task_id: str, task_history: list, reward: float): - """Get the trajectory from the task history.""" - pattern = r"\n\nSome Related Experience to help you to complete the task:.*" - task_history[1]["content"] = re.sub(pattern, "", task_history[1]["content"], flags=re.DOTALL) - return { - "task_id": task_id, - "messages": task_history, - "score": reward, - } - - def summary_memory(self, trajectories): - """Generate a summary of conversation messages and create task memories""" - - response = requests.post( - url=f"{self.memory_base_url}summary_task_memory", - json={ - "trajectories": trajectories, - "success_threshold": 1.0, - "enable_soft_comparison": True, - "validation_threshold": 0.5, - }, - ) - - result = self.handle_api_response(response) - if not result: - return [] - - # Extract memory list from response - memory_list = result.get("metadata", {}).get("memory_list", []) - print(f"Task memory list created: {len(memory_list)} memories") - return memory_list - - def add_memory(self, memory_list): - """Add the memory to the memory pool.""" - response = requests.post( - url=f"{self.memory_base_url}add_task_memory", - json={ - "memory_list": memory_list, - }, - ) - response.raise_for_status() - - def update_memory_information(self, memory_list, update_utility: bool = False): - """Update the memory information.""" - response = requests.post( - url=f"{self.memory_base_url}record_task_memory", - json={ - "memory_list": memory_list, - "update_utility": update_utility, - }, - ) - response.raise_for_status() - logger.info(response.json()) - - def delete_memory(self): - """Delete the memory from the memory pool.""" - response = requests.post( - url=f"{self.memory_base_url}delete_task_memory", - json={ - "freq_threshold": self.freq_threshold, - "utility_threshold": self.utility_threshold, - }, - ) - response.raise_for_status() - - -def main(): - """Main function to run the Appworld React Agent.""" - dataset_name = "train" - task_ids = load_task_ids(dataset_name) - agent = AppworldReactAgent(index=0, task_ids=task_ids[0:1], experiment_name=dataset_name, num_trials=1) - result = agent.execute() - logger.info(f"result={json.dumps(result)}") - - -if __name__ == "__main__": - main() diff --git a/benchmark/appworld/prompt.py b/benchmark/appworld/prompt.py deleted file mode 100644 index 45091967..00000000 --- a/benchmark/appworld/prompt.py +++ /dev/null @@ -1,660 +0,0 @@ -# flake8: noqa: E402, E501 -# pylint: disable=C0114,C0301 -# This is a basic prompt template containing all the necessary onboarding information to solve AppWorld tasks. It explains the role of the agent and the supervisor, how to explore the API documentation, how to operate the interactive coding environment and call APIs via a simple task, and provides key instructions and disclaimers. - -# You can adapt it as needed by your agent. You can also choose to bypass API docs app and build your own API retrieval, e.g., for FullCodeRefl, IPFunCall, etc, we asked an LLM to predict relevant APIs separately and put its documentation directly in the prompt. -PROMPT_TEMPLATE = """ -USER: -I am your supervisor and you are a super intelligent AI Assistant whose job is to achieve my day-to-day tasks completely autonomously. - -To do this, you will need to interact with app/s (e.g., spotify, venmo, etc) using their associated APIs on my behalf. For this you will undertake a *multi-step conversation* using a python REPL environment. That is, you will write the python code and the environment will execute it and show you the result, based on which, you will write python code for the next step and so on, until you've achieved the goal. This environment will let you interact with app/s using their associated APIs on my behalf. - -Here are three key APIs that you need to know to get more information - -# To get a list of apps that are available to you. -print(apis.api_docs.show_app_descriptions()) - -# To get the list of apis under any app listed above, e.g. supervisor -print(apis.api_docs.show_api_descriptions(app_name='supervisor')) - -# To get the specification of a particular api, e.g. supervisor app's show_account_passwords -print(apis.api_docs.show_api_doc(app_name='supervisor', api_name='show_account_passwords')) - -Each code execution will produce an output that you can use in subsequent calls. Using these APIs, you can now generate code, that the environment will execute, to solve the task. - -For example, consider the task: - -My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}. - -Task: - -What is the password for my Spotify account? - -ASSISTANT: -# Okay. Lets first find which apps are available to get the password by looking at the app descriptions. -print(apis.api_docs.show_app_descriptions()) - -USER: -[ - { - "name": "api_docs", - "description": "An app to search and explore API documentation." - }, - { - "name": "supervisor", - "description": "An app to access supervisor's personal information, account credentials, addresses, payment cards, and manage the assigned task." - }, - ... - { - "name": "spotify", - "description": "A music streaming app to stream songs and manage song, album and playlist libraries." - }, - { - "name": "venmo", - "description": "A social payment app to send, receive and request money to and from others." - }, - ... -] - - -ASSISTANT: -# Looks like the supervisor app could help me with that. Lets see what apis are available under this app. -print(apis.api_docs.show_api_descriptions(app_name='supervisor')) - - -USER: -[ - ... - "show_account_passwords : Show your supervisor's account passwords." - ... -] - - -ASSISTANT: -# I can use `show_account_passwords` to get the passwords. Let me see its detailed specification to understand its arguments and output structure. -print(apis.api_docs.show_api_doc(app_name='supervisor', api_name='show_account_passwords')) - -USER: -{ - 'app_name': 'supervisor', - 'api_name': 'show_account_passwords', - 'path': '/account_passwords', - 'method': 'GET', - 'description': "Show your supervisor's app account passwords.", - 'parameters': [], - 'response_schemas': { - 'success': [{'account_name': 'string', 'password': 'string'}], - 'failure': {'message': 'string'} - } -} - -ASSISTANT: -# Okay, it requires no arguments. So I can just call it directly. -print(apis.supervisor.show_account_passwords()) - -USER: -[ - { - "account_name": "spotify", - "password": "dummy_spotify_pass" - }, - { - "account_name": "file_system", - "password": "dummy_fs_pass" - }, - ... -] - - -ASSISTANT: -# So the Spotify password is an entry in the `passwords` list with the account_name=spotify. -spotify_password = [account_password["account_name"] == "spotify" for account_password in passwords][0]["password"] -print(spotify_password) - - -USER: -dummy_spotify_pass - -ASSISTANT: -# When the task is completed, I need to call apis.supervisor.complete_task(). If there is an answer, I need to pass it as an argument `answer`. I will pass the spotify_password as an answer. -apis.supervisor.complete_task(answer=spotify_password) - - -USER: -Marked the active task complete. - - ----------------------------------------------- - -USER: -**Key instructions and disclaimers**: - -1. The email addresses, access tokens and variables (e.g. spotify_password) in the example above were only for demonstration. Obtain the correct information by calling relevant APIs yourself. -2. Only generate valid code blocks, i.e., do not put them in ```...``` or add any extra formatting. Any thoughts should be put as code comments. -3. You can use the variables from the previous code blocks in the subsequent code blocks. -4. Write small chunks of code and only one chunk of code in every step. Make sure everything is working correctly before making any irreversible change. -5. The provided Python environment has access to its standard library. But modules and functions that have a risk of affecting the underlying OS, file system or process are disabled. You will get an error if do call them. -6. Any reference to a file system in the task instructions means the file system *app*, operable via given APIs, and not the actual file system the code is running on. So do not write code making calls to os-level modules and functions. -7. To interact with apps, only use the provided APIs, and not the corresponding Python packages. E.g., do NOT use `spotipy` for Spotify. Remember, the environment only has the standard library. -8. The provided API documentation has both the input arguments and the output JSON schemas. All calls to APIs and parsing its outputs must be as per this documentation. -9. For APIs that return results in "pages", make sure to consider all pages. -10. To obtain current date or time, use Python functions like `datetime.now()` or obtain it from the phone app. Do not rely on your existing knowledge of what the current date or time is. -11. For all temporal requests, use proper time boundaries, e.g., if I ask for something that happened yesterday, make sure to consider the time between 00:00:00 and 23:59:59. All requests are concerning a single, default (no) time zone. -12. Any reference to my friends, family or any other person or relation refers to the people in my phone's contacts list. -13. All my personal information, and information about my app account credentials, physical addresses and owned payment cards are stored in the "supervisor" app. You can access them via the APIs provided by the supervisor app. -14. Once you have completed the task, call `apis.supervisor.complete_task()`. If the task asks for some information, return it as the answer argument, i.e. call `apis.supervisor.complete_task(answer=)`. For tasks that do not require an answer, just skip the answer argument or pass it as None. -15. The answers, when given, should be just entity or number, not full sentences, e.g., `answer=10` for "How many songs are in the Spotify queue?". When an answer is a number, it should be in numbers, not in words, e.g., "10" and not "ten". -16. You can also pass `status="fail"` in the complete_task API if you are sure you cannot solve it and want to exit. -17. You must make all decisions completely autonomously and not ask for any clarifications or confirmations from me or anyone else. - -USER: -Using these APIs, now generate code to solve the actual task: - -My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}. - -Task: - -{{ instruction }} -""" - -PROMPT_TEMPLATE_WITH_EXPERIENCE = """ -USER: -I am your supervisor and you are a super intelligent AI Assistant whose job is to achieve my day-to-day tasks completely autonomously. - -To do this, you will need to interact with app/s (e.g., spotify, venmo, etc) using their associated APIs on my behalf. For this you will undertake a *multi-step conversation* using a python REPL environment. That is, you will write the python code and the environment will execute it and show you the result, based on which, you will write python code for the next step and so on, until you've achieved the goal. This environment will let you interact with app/s using their associated APIs on my behalf. - -Here are three key APIs that you need to know to get more information - -# To get a list of apps that are available to you. -print(apis.api_docs.show_app_descriptions()) - -# To get the list of apis under any app listed above, e.g. supervisor -print(apis.api_docs.show_api_descriptions(app_name='supervisor')) - -# To get the specification of a particular api, e.g. supervisor app's show_account_passwords -print(apis.api_docs.show_api_doc(app_name='supervisor', api_name='show_account_passwords')) - -Each code execution will produce an output that you can use in subsequent calls. Using these APIs, you can now generate code, that the environment will execute, to solve the task. - -For example, consider the task: - -My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}. - -Task: - -What is the password for my Spotify account? - -ASSISTANT: -# Okay. Lets first find which apps are available to get the password by looking at the app descriptions. -print(apis.api_docs.show_app_descriptions()) - -USER: -[ - { - "name": "api_docs", - "description": "An app to search and explore API documentation." - }, - { - "name": "supervisor", - "description": "An app to access supervisor's personal information, account credentials, addresses, payment cards, and manage the assigned task." - }, - ... - { - "name": "spotify", - "description": "A music streaming app to stream songs and manage song, album and playlist libraries." - }, - { - "name": "venmo", - "description": "A social payment app to send, receive and request money to and from others." - }, - ... -] - - -ASSISTANT: -# Looks like the supervisor app could help me with that. Lets see what apis are available under this app. -print(apis.api_docs.show_api_descriptions(app_name='supervisor')) - - -USER: -[ - ... - "show_account_passwords : Show your supervisor's account passwords." - ... -] - - -ASSISTANT: -# I can use `show_account_passwords` to get the passwords. Let me see its detailed specification to understand its arguments and output structure. -print(apis.api_docs.show_api_doc(app_name='supervisor', api_name='show_account_passwords')) - -USER: -{ - 'app_name': 'supervisor', - 'api_name': 'show_account_passwords', - 'path': '/account_passwords', - 'method': 'GET', - 'description': "Show your supervisor's app account passwords.", - 'parameters': [], - 'response_schemas': { - 'success': [{'account_name': 'string', 'password': 'string'}], - 'failure': {'message': 'string'} - } -} - -ASSISTANT: -# Okay, it requires no arguments. So I can just call it directly. -print(apis.supervisor.show_account_passwords()) - -USER: -[ - { - "account_name": "spotify", - "password": "dummy_spotify_pass" - }, - { - "account_name": "file_system", - "password": "dummy_fs_pass" - }, - ... -] - - -ASSISTANT: -# So the Spotify password is an entry in the `passwords` list with the account_name=spotify. -spotify_password = [account_password["account_name"] == "spotify" for account_password in passwords][0]["password"] -print(spotify_password) - - -USER: -dummy_spotify_pass - -ASSISTANT: -# When the task is completed, I need to call apis.supervisor.complete_task(). If there is an answer, I need to pass it as an argument `answer`. I will pass the spotify_password as an answer. -apis.supervisor.complete_task(answer=spotify_password) - - -USER: -Marked the active task complete. - - ----------------------------------------------- - -USER: -**Key instructions and disclaimers**: - -1. The email addresses, access tokens and variables (e.g. spotify_password) in the example above were only for demonstration. Obtain the correct information by calling relevant APIs yourself. -2. Only generate valid code blocks, i.e., do not put them in ```...``` or add any extra formatting. Any thoughts should be put as code comments. -3. You can use the variables from the previous code blocks in the subsequent code blocks. -4. Write small chunks of code and only one chunk of code in every step. Make sure everything is working correctly before making any irreversible change. -5. The provided Python environment has access to its standard library. But modules and functions that have a risk of affecting the underlying OS, file system or process are disabled. You will get an error if do call them. -6. Any reference to a file system in the task instructions means the file system *app*, operable via given APIs, and not the actual file system the code is running on. So do not write code making calls to os-level modules and functions. -7. To interact with apps, only use the provided APIs, and not the corresponding Python packages. E.g., do NOT use `spotipy` for Spotify. Remember, the environment only has the standard library. -8. The provided API documentation has both the input arguments and the output JSON schemas. All calls to APIs and parsing its outputs must be as per this documentation. -9. For APIs that return results in "pages", make sure to consider all pages. -10. To obtain current date or time, use Python functions like `datetime.now()` or obtain it from the phone app. Do not rely on your existing knowledge of what the current date or time is. -11. For all temporal requests, use proper time boundaries, e.g., if I ask for something that happened yesterday, make sure to consider the time between 00:00:00 and 23:59:59. All requests are concerning a single, default (no) time zone. -12. Any reference to my friends, family or any other person or relation refers to the people in my phone's contacts list. -13. All my personal information, and information about my app account credentials, physical addresses and owned payment cards are stored in the "supervisor" app. You can access them via the APIs provided by the supervisor app. -14. Once you have completed the task, call `apis.supervisor.complete_task()`. If the task asks for some information, return it as the answer argument, i.e. call `apis.supervisor.complete_task(answer=)`. For tasks that do not require an answer, just skip the answer argument or pass it as None. -15. The answers, when given, should be just entity or number, not full sentences, e.g., `answer=10` for "How many songs are in the Spotify queue?". When an answer is a number, it should be in numbers, not in words, e.g., "10" and not "ten". -16. You can also pass `status="fail"` in the complete_task API if you are sure you cannot solve it and want to exit. -17. You must make all decisions completely autonomously and not ask for any clarifications or confirmations from me or anyone else. -18. Some Related Experience to help you to complete the task: -{{experience}} - -USER: -Using these APIs, now generate code to solve the actual task: - -My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}. - -Task: - -{{ instruction }} -""" - -NEW_PROMPT_TEMPLATE = """ -USER: -I am your supervisor and you are a super intelligent AI Assistant whose job is to achieve my day-to-day tasks completely autonomously. - -To do this, you will need to interact with app/s (e.g., spotify, venmo etc) using their associated APIs on my behalf. For this you will undertake a *multi-step conversation* using a python REPL environment. That is, you will write the python code and the environment will execute it and show you the result, based on which, you will write python code for the next step and so on, until you've achieved the goal. This environment will let you interact with app/s using their associated APIs on my behalf. - -Here are three key APIs that you need to know to get more information - -# To get a list of apps that are available to you. - -```python -print(apis.api_docs.show_app_descriptions()) -``` - -# To get the list of apis under any app listed above, e.g. spotify - -```python -print(apis.api_docs.show_api_descriptions(app_name='spotify')) -``` - -# To get the specification of a particular api, e.g. spotify app's login api - -```python -print(apis.api_docs.show_api_doc(app_name='spotify', api_name='login')) -``` - -Each code execution will produce an output that you can use in subsequent calls. Using these APIs, you can now generate code, that I will execute, to solve the task. Let's start with the task - -My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}. -Task: How many playlists do I have in Spotify? - -ASSISTANT: -Okay. Lets first find which APIs are available to use in Spotify. -Code: -```python -print(apis.api_docs.show_api_descriptions(app_name='spotify')) -``` - -USER: -Output: -``` -[ - ... - "login : Login to your account.", - "logout : Logout from your account.", - ... -] -``` - - -ASSISTANT: -Okay. Looks like I can use the `login` api. Lets find its specifications. - -Code: -```python -print(apis.api_docs.show_api_doc(app_name='spotify', api_name='login')) -``` - - -USER: -Output: -``` -{ - "app_name": "spotify", - "api_name": "login", - "path": "/auth/token", - "method": "POST", - "description": "Login to your account.", - "parameters": [ - { - "name": "username", - "type": "string", - "required": true, - "description": "Your account email.", - "default": null, - "constraints": [] - }, - { - "name": "password", - "type": "string", - "required": true, - "description": "Your account password.", - "default": null, - "constraints": [] - } - ], - "response_schemas": { - "success": { - "token_type": "string", - "access_token": "string" - }, - "failure": { - "message": "string" - } - } -} -``` - - -ASSISTANT: -I need the supervisor's username and password. Lets see if any app can help me get that. - -Code: -```python -print(apis.api_docs.show_app_descriptions()) -``` - - -USER: -Output: -``` -{{ app_descriptions }} -``` - - -ASSISTANT: -Looks like the supervisor app could help me with that. - -Code: -```python -print(apis.api_docs.show_api_descriptions(app_name='supervisor')) -``` - - -USER: -Output: -``` -[ - ... - "show_account_passwords : Show your supervisor's account passwords." - ... -] -``` - - -ASSISTANT: -I can use `show_account_passwords` to get the passwords. So I will call that and save it in a variable for future reference. -Code: -```python -passwords=apis.supervisor.show_account_passwords() -print(passwords) -``` - - -USER: -Output: -``` -[ - { - "account_name": "spotify", - "password": "dummy_spotify_pass" - }, - { - "account_name": "file_system", - "password": "dummy_fs_pass" - } -] -``` - - -ASSISTANT: -So the Spotify password is an entry in the `passwords` list with the account_name=spotify. I can use it to login now with the provided email address: {{ supervisor.email }}. - -Code: - ```python -spotify_password = [account_password["account_name"] == "spotify" for account_password in passwords][0]["password"] -login_result = apis.spotify.login(username='{{ supervisor.email }}', password=spotify_password) -print(login_result) -``` - - -USER: -Output: -``` -{ - "token_type": "Bearer", - "access_token": "fake_access_token" -} -``` - - -ASSISTANT: -First, let's check the available APIs for the Spotify app. -Code: -```python -print(apis.api_docs.show_api_descriptions(app_name='spotify')) -``` - - -USER: -Output: -``` -[ -... -{ - "name": "show_playlist_library", - "description": "Get a list of playlists in the user's playlist library." -}, -{ - "name": "show_playlist", - "description": "Get detailed information about a specific playlist. You can view your own playlists or others' playlists if they are public." - }, -... -] -``` - - -ASSISTANT: -The `show_playlist_library` API seems to be the one we need. Let's check its specifications. -Code: -```python -print(apis.api_docs.show_api_doc(app_name='spotify', api_name='show_playlist_library')) -``` - - -USER: -Output: -``` -{ - "app_name": "spotify", - "api_name": "show_playlist_library", - "path": "/private_playlists", - "method": "GET", - "description": "Get a list of playlists in the user's playlist library.", - "parameters": [ - { - "name": "access_token", - "type": "string", - "required": true, - "description": "Access token obtained from spotify app login.", - "default": null, - "constraints": [] - }, - { - "name": "page_index", - "type": "integer", - "required": false, - "description": "The index of the page to retrieve.", - "default": 0, - "constraints": [ - "value >= 0.0" - ] - }, - { - "name": "page_limit", - "type": "integer", - "required": false, - "description": "The maximum number of results to return per page.", - "default": 5, - "constraints": [ - "value >= 1.0, <= 20.0" - ] - }, - { - "name": "is_public", - "type": "boolean", - "required": false, - "description": "Whether to show public playlists or private playlists.", - "default": null, - "constraints": [] - } - ], - "response_schema": [ - { - "title": "string", - "created_at": "2019-01-01T00:00:00", - "is_public": true, - "rating": 0.0, - "like_count": 1, - "owner_email": "user@example.com", - "playlist_id": 1, - "song_ids": [ - 1 - ] - } - ] -} -``` - - -ASSISTANT: -I need to page through all the playlists to get the list of playlists and save it in `playlists`. -Code: -```python -page_index = 0 -playlists = [] -while page_index < 10: - playlist_page = apis.spotify.show_playlist_library(access_token=spotify_access_token, page_index=page_index) - if playlist_page: - playlists.extend(playlist_page) - page_index += 1 - else: - break -num_playlists = len(playlists) -print(num_playlists) - -``` - -USER: -Output: -``` -23 -``` - - -ASSISTANT: -Now that the task is completed, I can call apis.supervisor.complete_task(). Since this task has an answer to be returned, I will pass that as an argument. - -Code: -```python -apis.supervisor.complete_task(answer=num_playlists) -``` - - -USER: -Output: -Marked the active task complete. - - ----------------------------------------------- - -USER: -**Key instructions**: -(1) Make sure to end code blocks with ``` followed by a newline(\n). - -(2) Remember you can use the variables in your code in subsequent code blocks. - -(3) Remember that the email addresses, access tokens and variables (e.g. spotify_password) in the example above are not valid anymore. - -(4) You can use the "supervisor" app to get information about my accounts and use the "phone" app to get information about friends and family. - -(5) Always look at API specifications (using apis.api_docs.show_api_doc) before calling an API. - -(6) Write small chunks of code and only one chunk of code in every step. Make sure everything is working correctly before making any irreversible change. - -(7) Many APIs return items in "pages". Make sure to run through all the pages by looping over `page_index`. - -(8) Once you have completed the task, make sure to call apis.supervisor.complete_task(). If the task asked for some information, return it as the answer argument, i.e. call apis.supervisor.complete_task(answer=). Many tasks do not require an answer, so in those cases, just call apis.supervisor.complete_task() i.e. do not pass any argument. - -USER: -Using these APIs, now generate code to solve the actual task: - -My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}. - -""" diff --git a/benchmark/appworld/run_appworld.py b/benchmark/appworld/run_appworld.py deleted file mode 100644 index b0cde1df..00000000 --- a/benchmark/appworld/run_appworld.py +++ /dev/null @@ -1,197 +0,0 @@ -# pylint: disable=E0611 -"""Run the Appworld React Agent.""" - -import os -import json -import time -from pathlib import Path - -import ray -import requests -from loguru import logger -from dotenv import load_dotenv -from appworld import load_task_ids -from appworld_react_agent import AppworldReactAgent - -os.environ["APPWORLD_ROOT"] = "." - -load_dotenv("../../.env") - - -def run_agent( - run_index: int, - max_workers: int, - model_name: str, - dataset_name: str, - experiment_suffix: str, - num_trials: int = 1, - use_memory: bool = False, - memory_base_url: str = "http://0.0.0.0:8002/", - use_memory_addition: bool = False, - use_memory_deletion: bool = False, - delete_freq: int = 10, - freq_threshold: int = 5, - utility_threshold: float = 0.5, - batch_size: int = 4, -): - """Run the Appworld React Agent.""" - experiment_name = dataset_name + "_" + experiment_suffix - path: Path = Path(f"./exp_result/{model_name}") - path.mkdir(parents=True, exist_ok=True) - - task_ids = load_task_ids(dataset_name) - - result: list = [] - - def dump_file(): - with open(path / f"{experiment_name}.jsonl", "a", encoding="utf-8") as f: - for x in result: - f.write(json.dumps(x) + "\n") - - if max_workers > 1: - # Process tasks in batches - total_tasks = len(task_ids) - num_batches = (total_tasks + batch_size - 1) // batch_size # Ceiling division - - logger.info(f"Total tasks: {total_tasks}, Batch size: {batch_size}, Number of batches: {num_batches}") - - for batch_idx in range(num_batches): - # Initialize Ray for this batch - start_idx = batch_idx * batch_size - end_idx = min(start_idx + batch_size, total_tasks) - batch_task_ids = task_ids[start_idx:end_idx] - - logger.info(f"Starting batch {batch_idx + 1}/{num_batches} with {len(batch_task_ids)} tasks") - - # Initialize Ray with the number of CPUs needed for this batch - ray.init(num_cpus=len(batch_task_ids)) - - future_list: list = [] - for i, task_id in enumerate(batch_task_ids): - actor = AppworldReactAgent.remote( - index=start_idx + i, - model_name=model_name, - task_ids=[task_id], - experiment_name=experiment_name, - num_trials=num_trials, - use_memory=use_memory, - memory_base_url=memory_base_url, - use_memory_addition=use_memory_addition, - use_memory_deletion=use_memory_deletion, - delete_freq=delete_freq, - freq_threshold=freq_threshold, - utility_threshold=utility_threshold, - ) - future = actor.execute.remote() - future_list.append(future) - time.sleep(1) - - logger.info(f"Batch {batch_idx + 1} submit complete, waiting for results...") - - # Collect results from this batch - for i, (task_id, future) in enumerate(zip(batch_task_ids, future_list)): - try: - t_result = ray.get(future) - if t_result: - if isinstance(t_result, list): - result.extend(t_result) - else: - result.append(t_result) - except Exception: - logger.exception(f"run ray error with task_id={task_id}") - - logger.info(f"Batch {batch_idx + 1}: task {i + 1}/{len(batch_task_ids)} complete") - - # Shutdown Ray to free resources before next batch - ray.shutdown() - logger.info(f"Batch {batch_idx + 1}/{num_batches} complete, Ray resources released") - - # Optional: small delay between batches - if batch_idx < num_batches - 1: - time.sleep(2) - - dump_file() - - else: - agent = AppworldReactAgent( - index=run_index, - model_name=model_name, - task_ids=task_ids, - experiment_name=experiment_name, - num_trials=num_trials, - use_memory=use_memory, - memory_base_url=memory_base_url, - use_memory_addition=use_memory_addition, - use_memory_deletion=use_memory_deletion, - delete_freq=delete_freq, - freq_threshold=freq_threshold, - utility_threshold=utility_threshold, - ) - result = agent.execute() - - dump_file() - - -def handle_api_response(response: requests.Response): - """Handle API response with proper error checking""" - if response.status_code != 200: - print(f"Error: {response.status_code}") - print(response.text) - return None - - return response.json() - - -def load_memory(path: str = "docs/library", api_url: str = "http://0.0.0.0:8002/"): - """Load memories from disk into the vector store""" - response = requests.post( - url=f"{api_url}load_memory", - json={ - "load_file_path": path, - "clear_existing": True, - }, - ) - - result = handle_api_response(response) - if result: - print(f"Memory loaded from {path}") - - -def main(): - """Main function to run the Appworld React Agent.""" - max_workers = 16 - batch_size = 8 - - num_runs = 4 # Number of runs - num_trials = 1 # for self-reflection - model_name = "qwen3-8b" - use_memory = True - use_memory_addition = False - use_memory_deletion = False - memory_base_url = "http://0.0.0.0:8002/" - - if use_memory: - load_file_path = "docs/library/paper_data/task/appworld_qwen3_8b.jsonl" - load_memory(load_file_path, memory_base_url) - - for i in range(num_runs): - run_agent( - run_index=i, - max_workers=max_workers, - model_name=model_name, - dataset_name="test_normal", - experiment_suffix="with-fixed-memory", - num_trials=num_trials, - use_memory=use_memory, - memory_base_url=memory_base_url, - use_memory_addition=use_memory_addition, - use_memory_deletion=use_memory_deletion, - delete_freq=5, - freq_threshold=5, - utility_threshold=0.5, - batch_size=batch_size, - ) - - -if __name__ == "__main__": - main() diff --git a/benchmark/appworld/run_exp_statistic.py b/benchmark/appworld/run_exp_statistic.py deleted file mode 100644 index 9e45effb..00000000 --- a/benchmark/appworld/run_exp_statistic.py +++ /dev/null @@ -1,164 +0,0 @@ -"""Run the experiment statistic.""" - -import json -from collections import defaultdict -from pathlib import Path - -import pandas as pd -from loguru import logger - - -def calculate_best_at_k(scores: list, k: int) -> float: - """ - Calculate best@k - Divide scores into groups of size k, take the maximum value in each group, - then average these maximum values - - Args: - scores: List of after_score values for all runs of a task - k: Group size - - Returns: - best@k value - """ - if len(scores) % k != 0: - raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})") - - group_maxs = [] - for i in range(0, len(scores), k): - group = scores[i : i + k] - group_maxs.append(max(group)) - - return sum(group_maxs) / len(group_maxs) - - -def calculate_pass_at_k(scores: list, k: int) -> float: - """Calculate pass@k.""" - if len(scores) % k != 0: - raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})") - - group_maxs = [] - for i in range(0, len(scores), k): - group = scores[i : i + k] - is_pass = 1.0 if max(group) >= 1.0 else 0.0 - group_maxs.append(is_pass) - - return sum(group_maxs) / len(group_maxs) - - -def get_possible_k_values(total_runs: int) -> list: - """ - Get all possible k values (factors of total_runs) - - Args: - total_runs: Total number of runs - - Returns: - List of k values in descending order - """ - k_values = [] - for k in range(1, total_runs + 1): - if total_runs % k == 0: - k_values.append(k) - return sorted(k_values, reverse=True) # Sort from large to small - - -def run_exp_statistic(): - """Run the experiment statistic.""" - path: Path = Path("./exp_result/qwen3-8b") - - # Store results for all experiments - all_results = {} - - for file in path.glob("*.jsonl"): # [f for f in path.glob("*.jsonl") if not f.stem[-1].isdigit()] - # Group results by task_id - task_results = defaultdict(list) - - with open(file, "r", encoding="utf-8") as f: - for line in f: - if not line.strip(): - continue - data = json.loads(line) - - if isinstance(data, list): - for part_data in data: - task_id = part_data["task_id"] - after_score = part_data["after_score"] - task_results[task_id].append(after_score) - else: - task_id = data["task_id"] - after_score = data["after_score"] - task_results[task_id].append(after_score) - - if not task_results: - logger.warning(f"No valid data found in file {file}") - continue - - # Check if each task has consistent number of runs - run_counts = [len(scores) for scores in task_results.values()] - if len(set(run_counts)) > 1: - logger.warning(f"Inconsistent number of runs for different tasks in file {file}: {set(run_counts)}") - continue - - num_runs = run_counts[0] - logger.info(f"File {file}: {len(task_results)} tasks, {num_runs} runs per task") - - # Get all possible k values - k_values = get_possible_k_values(num_runs) - logger.info(f"Calculable best@k values: {k_values}") - - # Calculate various best@k values - file_results = {"file": file.name} - - for k in k_values: - best_at_k_scores = [] - pass_at_k_scores = [] - for task_id, scores in task_results.items(): - try: - best_k_score = calculate_best_at_k(scores, k) - pass_at_k_score = calculate_pass_at_k(scores, k) - pass_at_k_scores.append(pass_at_k_score) - best_at_k_scores.append(best_k_score) - except ValueError as e: - logger.error(f"Error calculating best@{k} for task {task_id}: {e}") - continue - - if best_at_k_scores: - avg_best_at_k = sum(best_at_k_scores) / len(best_at_k_scores) - file_results[f"best@{k}"] = avg_best_at_k - logger.info(f"file={file.name} best@{k}={avg_best_at_k:.4f}") - - if pass_at_k_scores: - avg_pass_at_k = sum(pass_at_k_scores) / len(pass_at_k_scores) - file_results[f"pass@{k}"] = avg_pass_at_k - logger.info(f"file={file.name} pass@{k}={avg_pass_at_k:.4f}") - - all_results[file.name] = file_results - - # Create and display table - if all_results: - df = pd.DataFrame(list(all_results.values())) - df = df.set_index("file") - - # Sort columns by the number in column name (best@8, best@4, best@2, best@1) - pass_columns = [col for col in df.columns if col.startswith("pass@")] - # best_columns = [col for col in df.columns] - pass_columns.sort(key=lambda x: x, reverse=False) - df = df[pass_columns] - - print("\n" + "=" * 80) - print("Experiment Results Summary Table") - print("=" * 80) - print(df.round(4)) - print("=" * 80) - - # Save table to CSV - output_path = path / "experiment_summary.csv" - df.to_csv(output_path) - logger.info(f"Results table saved to: {output_path}") - else: - logger.warning("No valid experiment results found") - - -if __name__ == "__main__": - run_exp_statistic() diff --git a/benchmark/bfcl/bfcl_agent.py b/benchmark/bfcl/bfcl_agent.py deleted file mode 100644 index 4e1e12b7..00000000 --- a/benchmark/bfcl/bfcl_agent.py +++ /dev/null @@ -1,726 +0,0 @@ -# flake8: noqa: E402 -# pylint: disable=too-many-return-statements -"""A minimal ReAct Agent for BFCL-v3(multi-turn) tasks.""" - -import re -import os -import time -import json -import warnings -import tempfile -import datetime -from pathlib import Path -from typing import Dict, List, Any - -import ray -import requests -from tqdm import tqdm -from loguru import logger -from openai import OpenAI -from dotenv import load_dotenv - -from bfcl_utils import ( - load_test_case, - handle_user_turn, - handle_tool_calls, - extract_tool_schema, - extract_single_turn_response, - extract_multi_turn_responses, - capture_and_print_score_files, - create_error_response, -) -from bfcl_eval.model_handler.api_inference.qwen import QwenAPIHandler -from bfcl_eval.eval_checker.multi_turn_eval.multi_turn_utils import ( - is_empty_execute_response, -) -from bfcl_eval.eval_checker.eval_runner import ( - multi_turn_runner, - ast_file_runner, -) -from bfcl_eval.eval_checker.eval_runner_helper import record_cost_latency -from bfcl_eval.utils import ( - is_multi_turn, - is_relevance_or_irrelevance, - find_file_with_suffix, - load_file, -) - -os.environ["BFCL_DATA_PATH"] = "data/multiturn_data_base_val.jsonl" -os.environ["BFCL_ANSWER_PATH"] = "data/possible_answer" -load_dotenv("../../.env") - - -@ray.remote -class BFCLAgent: - """A minimal ReAct Agent for BFCL-v3(multi-turn) tasks.""" - - def __init__( - self, - index: int, - task_ids: List[str], - experiment_name: str, - data_path: str = os.getenv("BFCL_DATA_PATH"), - answer_path: Path = Path(os.getenv("BFCL_ANSWER_PATH")), - model_name: str = "qwen3-8b", - temperature: float = 0.9, - max_interactions: int = 30, - max_response_size: int = 2000, - num_trials: int = 1, - enable_thinking: bool = False, - use_memory: bool = False, - use_memory_addition: bool = False, - use_memory_deletion: bool = False, - delete_freq: int = 10, - freq_threshold: int = 5, - utility_threshold: float = 0.5, - memory_base_url: str = "http://0.0.0.0:8002/", - ): - - self.index: int = index - self.task_ids: List[str] = task_ids - self.categories: List[str] = [task_id.rsplit("_", 1)[0] if "_" in task_id else task_id for task_id in task_ids] - self.experiment_name: str = experiment_name - self.data_path: str = data_path - self.answer_path: Path = answer_path - self.model_name: str = model_name - self.temperature: float = temperature - self.max_interactions: int = max_interactions - self.max_response_size: int = max_response_size - self.num_trials: int = num_trials - self.enable_thinking: bool = enable_thinking - self.use_memory: bool = use_memory - self.use_memory_addition: bool = use_memory_addition if use_memory else False - self.use_memory_deletion: bool = use_memory_deletion if use_memory else False - self.delete_freq: int = delete_freq - self.freq_threshold: int = freq_threshold - self.utility_threshold: float = utility_threshold - self.memory_base_url: str = memory_base_url - - self.llm_client = OpenAI() - - self.history: List[List[List[dict]]] = [[] for _ in range(num_trials)] - self.retrieved_memory_list: List[List[List[Any]]] = [[] for _ in range(num_trials)] - self.test_entry: List[List[Dict[str, Any]]] = [[] for _ in range(num_trials)] - self.original_test_entry: List[List[Dict[str, Any]]] = [[] for _ in range(num_trials)] - self.tool_schema: List[List[List[dict]]] = [[] for _ in range(num_trials)] - self.current_turn = [[0 for _ in range(len(task_ids))] for _ in range(num_trials)] - - for run_id in range(num_trials): - for task_index in range(len(task_ids)): - self.init_state(run_id, task_index) - - def init_state(self, run_id, i) -> Dict[str, Any]: - """Initialize the state of the agent.""" - self.test_entry[run_id].append(load_test_case(self.data_path, self.task_ids[i])) - self.original_test_entry[run_id].append(self.test_entry[run_id][i].get("extra", {})) - self.tool_schema[run_id].append(extract_tool_schema(self.test_entry[run_id][i].get("tools", [{}]))) - - msg = self.test_entry[run_id][i].get("messages", []) - self.history[run_id].append(msg) - self.retrieved_memory_list[run_id].append([]) - self.current_turn[run_id][i] = 1 - - def update_task_history_with_memory(self, run_id, task_index, previous_memories: None): - """Update the task history with memory.""" - query = self.history[run_id][task_index][0]["content"] - if len(previous_memories) == 0: - response = self.get_memory(query) - if response and "memory_list" in response["metadata"]: - self.retrieved_memory_list[run_id][task_index] = response["metadata"]["memory_list"] - task_memory = re.sub(r"\bMemory\s*(\d+)\s*[:]", r"Experience \1 :", response["answer"]) - logger.info(f"loaded task_memory: {task_memory}") - self.history[run_id][task_index][0] = self.get_query_with_memory(query, task_memory) - else: - formatted_memories = [] - for i, memory in enumerate(previous_memories, 1): - condition = memory["when_to_use"] - memory_content = memory["content"] - memory_text = f"Experience {i} :\n When to use: {condition}\n Content: {memory_content}\n" - formatted_memories.append(memory_text) - self.history[run_id][task_index][0] = self.get_query_with_memory(query, "\n".join(formatted_memories)) - - def get_query_with_memory(self, query: str, memory: str): - """Get the query with memory.""" - return { - "role": "user", - "content": "Task:\n" + query + "\n\nSome Related Experience to help you to complete the task:\n" + memory, - } - - def get_query_without_experience(self, query: str): - """Get the query without experience.""" - if "\n\nSome Related Experience" in query: - query = query.split("\n\nSome Related Experience")[0].split("Task:\n")[-1] - return query - - def get_traj_from_task_history(self, task_id: str, task_history: list, reward: float): - """Get the trajectory from the task history.""" - return { - "task_id": task_id, - "messages": task_history, - "score": reward, - } - - def handle_api_response(self, response: requests.Response): - """Handle API response with proper error checking""" - if response.status_code != 200: - print(f"Error: {response.status_code}") - print(response.text) - return None - - return response.json() - - def get_memory(self, query: str): - """Retrieve relevant task memories based on a query""" - response = requests.post( - url=f"{self.memory_base_url}retrieve_task_memory", - json={ - "query": query, - "enable_llm_rerank": False, - "enable_score_filter": False, - "top_k": 5, - "enable_llm_rewrite": False, - }, - ) - - result = self.handle_api_response(response) - if not result: - return None - - logger.info(f"query: {query}, response: {result}") - return result - - def summary_memory(self, trajectories): - """Generate a summary of conversation messages and create task memories""" - response = requests.post( - url=f"{self.memory_base_url}summary_task_memory", - json={ - "trajectories": trajectories, - "success_threshold": 1.0, - "enable_soft_comparison": True, - "validation_threshold": 0.5, - }, - ) - - result = self.handle_api_response(response) - if not result: - return [] - - # Extract memory list from response - memory_list = result.get("metadata", {}).get("memory_list", []) - logger.info(f"add new memories: {memory_list}") - return memory_list - - def add_memory(self, memory_list): - """Add the memory to the memory pool.""" - response = requests.post( - url=f"{self.memory_base_url}add_task_memory", - json={ - "memory_list": memory_list, - }, - ) - response.raise_for_status() - - def update_memory_information(self, memory_list, update_utility: bool = False): - """Update the memory information.""" - response = requests.post( - url=f"{self.memory_base_url}record_task_memory", - json={ - "memory_list": memory_list, - "update_utility": update_utility, - }, - ) - response.raise_for_status() - logger.info(response.json()) - - def delete_memory(self): - """Delete the memory from the memory pool.""" - response = requests.post( - url=f"{self.memory_base_url}delete_task_memory", - json={ - "freq_threshold": self.freq_threshold, - "utility_threshold": self.utility_threshold, - }, - ) - response.raise_for_status() - - def call_llm(self, messages: list, tool_schemas: list[dict]) -> str: - """Call the LLM.""" - for i in range(100): - try: - response = self.llm_client.chat.completions.create( - model=self.model_name, - messages=messages, - tools=tool_schemas, - temperature=self.temperature, - seed=0, - extra_body={"enable_thinking": self.enable_thinking}, - stream=self.enable_thinking, - parallel_tool_calls=True, - ) - if not self.enable_thinking: - out_msg = response.choices[0].message - return out_msg.model_dump(exclude_unset=True, exclude_none=True) - else: - reasoning_content = "" # Complete reasoning process - answer_content = "" # Define complete response - tool_info = [] # Store tool invocation information - is_answering = ( - False # Determine whether the reasoning process has finished and response has started - ) - - for chunk in response: - if not chunk.choices: - # Handle usage information - continue - - delta = chunk.choices[0].delta - # Handle AI's thought process (chain reasoning) - if hasattr(delta, "reasoning_content") and delta.reasoning_content is not None: - reasoning_content += delta.reasoning_content - # Handle final response content - else: - if not is_answering: # Print title when entering the response phase for the first time - is_answering = True - if delta.content is not None: - answer_content += delta.content - - # Handle tool invocation information (support parallel tool calls) - if delta.tool_calls is not None: - for tool_call in delta.tool_calls: - index = tool_call.index # Tool call index, used for parallel calls - - # Dynamically expand tool information storage list - while len(tool_info) <= index: - tool_info.append( - { - "id": "", - "type": "function", - "index": index, - "function": {"name": "", "arguments": ""}, - }, - ) - - # Collect tool call ID (used for subsequent function calls) - if tool_call.id: - tool_info[index]["id"] += tool_call.id - - # Collect function name (used for subsequent routing to specific functions) - if tool_call.function and tool_call.function.name: - tool_info[index]["function"]["name"] += tool_call.function.name - - # Collect function parameters (in JSON string format, need subsequent parsing) - if tool_call.function and tool_call.function.arguments: - tool_info[index]["function"]["arguments"] += tool_call.function.arguments - msg = { - "role": "assistant", - "content": answer_content, - "reasoning_content": reasoning_content, - } - if tool_info: - msg["tool_calls"] = tool_info - return msg - except Exception as e: - logger.exception(f"encounter error with {e.args}") - time.sleep(1 + i * 10) - - return "call llm error" - - def env_step(self, run_id: int, index: int, messages: str) -> str: - """ - Process one step in the conversation. - Both single turn and multi turn are supported. - - Args: - messages: List of conversation messages, with the last one being assistant response - test_entry: Test entry containing initial_config, involved_classes, question etc. - **kwargs: Additional arguments for compatibility - - Returns: - Dict containing next message and tools if applicable - """ - try: - if not messages: - return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index]) - - if messages[-1]["role"] != "assistant": - return create_error_response( - "Last message must be from assistant", - ) - - if "tool_calls" in messages[-1] and len(messages[-1]["tool_calls"]) > 0: - try: - tool_calls = messages[-1]["tool_calls"] - decoded_calls = self._convert_tool_calls_to_execution_format( - tool_calls, - ) - # decoded_calls:[function(param=xxx)] - print(f"decoded_calls: {decoded_calls}") - if is_empty_execute_response(decoded_calls): - warnings.warn( - f"is_empty_execute_response: {is_empty_execute_response(decoded_calls)}", - ) - return handle_user_turn( - self.original_test_entry[run_id][index], - self.current_turn[run_id][index], - ) - return handle_tool_calls( - tool_calls, - decoded_calls, - self.original_test_entry[run_id][index], - self.current_turn[run_id][index], - ) - except Exception as e: - warnings.warn(f"Errors during tool invocation: {str(e)}") - return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index]) - else: - return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index]) - - except Exception as e: - return create_error_response(f"Failed to process request: {str(e)}") - - def _convert_tool_calls_to_execution_format( - self, - tool_calls: List[Dict[str, Any]], - ) -> List[str]: - """ - Convert OpenAI format tool calls to execution format. - - Args: - tool_calls: List of tool calls in OpenAI format - - Returns: - List of function calls in string format - """ - execution_list = [] - - for tool_call in tool_calls: - function = tool_call.get("function", {}) - function_name = function.get("name", "") - - try: - arguments = function.get("arguments", "{}") - if isinstance(arguments, str): - args_dict = json.loads(arguments) - else: - args_dict = arguments - - args_str = ", ".join([f"{k}={repr(v)}" for k, v in args_dict.items()]) - execution_list.append(f"{function_name}({args_str})") - - except Exception: - execution_list.append(f"{function_name}()") - - return execution_list - - def get_reward(self, run_id, index) -> float: - """Get the reward.""" - try: - if not self.history[run_id][index] or not self.original_test_entry[run_id][index]: - return 0.0 - - model_name = "env_handler" - handler = QwenAPIHandler( - model_name, - temperature=1.0, - ) # FIXME: magic number - - model_result_data = self._convert_conversation_to_eval_format(run_id, index) - - prompt_data = [self.original_test_entry[run_id][index]] - - state = {"leaderboard_table": {}} - record_cost_latency( - state["leaderboard_table"], - model_name, - [model_result_data], - ) - - if is_relevance_or_irrelevance(self.categories[index]): - accuracy, _ = self._eval_relevance_test( - handler, - model_result_data, - prompt_data, - model_name, - self.category, - ) - else: - # Find the corresponding possible answer file - - possible_answer_file = find_file_with_suffix( - self.answer_path, - self.categories[index], - ) - possible_answer = load_file(possible_answer_file, sort_by_id=True) - possible_answer = [item for item in possible_answer if item["id"] == self.task_ids[index]] - if is_multi_turn(self.categories[index]): - accuracy, _ = self._eval_multi_turn_test( - handler, - model_result_data, - prompt_data, - possible_answer, - model_name, - self.categories[index], - ) - else: - accuracy, _ = self._eval_single_turn_test( - handler, - model_result_data, - prompt_data, - possible_answer, - model_name, - self.categories[index], - ) - print(f"model_result_data: {model_result_data}") - if possible_answer: - print(f"possible_answer: {possible_answer}") - else: - print("possible_answer: None") - - return accuracy - - except Exception: - import traceback - - traceback.print_exc() - return 0 - - def _convert_conversation_to_eval_format(self, run_id, index) -> Dict[str, Any]: - """ - Convert conversation history to evaluation format. - - Args: - conversation_result: Result from run_conversation - original_test_entry: Original test entry data - - Returns: - Data in format expected by multi_turn_runner or other runners - """ - if is_multi_turn(self.categories[index]): - turns_data = extract_multi_turn_responses(self.history[run_id][index]) - else: - turns_data = extract_single_turn_response(self.history[run_id][index]) - - model_result_data = { - "id": self.task_ids[index], - "result": turns_data, - "latency": 0, - "input_token_count": 0, - "output_token_count": 0, - } - - return model_result_data - - def _eval_multi_turn_test( - self, - handler, - model_result_data, - prompt_data, - possible_answer, - model_name, - test_category, - ): - """ - Evaluate multi-turn test. - - Args: - handler: Model handler instance - model_result_data: Model result data - prompt_data: Prompt data - possible_answer: Possible answer data - model_name: Name of the model - test_category: Category of the test - - Returns: - Tuple of (accuracy, total_count) - """ - with tempfile.TemporaryDirectory() as temp_dir: - score_dir = Path(temp_dir) - accuracy, total_count = multi_turn_runner( - handler=handler, - model_result=[model_result_data], - prompt=prompt_data, - possible_answer=possible_answer, - model_name=model_name, - test_category=test_category, - score_dir=score_dir, - ) - capture_and_print_score_files( - score_dir, - model_name, - test_category, - "multi_turn", - ) - return accuracy, total_count - - def _eval_single_turn_test( - self, - handler, - model_result_data, - prompt_data, - possible_answer, - model_name, - test_category, - ): - """ - Evaluate single-turn AST test. - - Args: - handler: Model handler instance - model_result_data: Model result data - prompt_data: Prompt data - possible_answer: Possible answer data - model_name: Name of the model - test_category: Category of the test - - Returns: - Tuple of (accuracy, total_count) - """ - language = "Python" - if "java" in test_category.lower(): - language = "Java" - elif "js" in test_category.lower() or "javascript" in test_category.lower(): - language = "JavaScript" - - with tempfile.TemporaryDirectory() as temp_dir: - score_dir = Path(temp_dir) - accuracy, total_count = ast_file_runner( - handler=handler, - model_result=[model_result_data], - prompt=prompt_data, - possible_answer=possible_answer, - language=language, - test_category=test_category, - model_name=model_name, - score_dir=score_dir, - ) - capture_and_print_score_files( - score_dir, - model_name, - test_category, - "single_turn", - ) - return accuracy, total_count - - def execute(self): - """Execute the agent.""" - result = [] - counter = 0 - for task_index, task_id in enumerate(tqdm(self.task_ids, desc=f"ray_index={self.index}")): - t_result = None - previous_memories = [] - for run_id in range(self.num_trials): - try: - start_time = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") - for i in range(self.max_interactions): - if self.use_memory and i == 0: - self.update_task_history_with_memory(run_id, task_index, previous_memories) - llm_output = self.call_llm( - self.history[run_id][task_index], - self.tool_schema[run_id][task_index], - ) - self.history[run_id][task_index].append(llm_output) - - env_output = self.env_step(run_id, task_index, self.history[run_id][task_index]) - # Possible env_output returns after environment interaction: - # 1. Triggers a query with available tools list: - # {"messages": [{"role": "user", "content": user_query}], "tools": tools} - # 2. Returns tool invocation result: {"messages": - # [{"role": "tool", "content": {}, 'tool_call_id': 'chatcmpl-tool-xxx'}]} - # : when success, returns result dicts, e.g., {"travel_cost_list": [x]}, - # when error, returns error message, - # e.g., {"error": "cd: temporary: No such directory. You cannot use path ..."} - # 3. Conversation completion: - # {"messages": [{"role": "env", "content": "[CONVERSATION_COMPLETED]"}]} - # 4. Program error: {"messages": [{"role": "env", "content": f"[ERROR] {error_message}"}]} - - # tool_list update - if "tools" in env_output: - self.tool_schema[run_id][task_index] = extract_tool_schema(env_output["tools"]) - - new_tool_calls = [] - new_tool_call_ids = [] - next_user_msg = "" - for idx, msg in enumerate(env_output.get("messages", [])): - if msg["role"] == "tool" and len(msg["content"]) > 0: - new_tool_calls.append(msg.get("content", "")) - new_tool_call_ids.append(msg.get("tool_call_id", "")) - elif msg["role"] == "user": - next_user_msg = msg.get("content", "") - self.current_turn[run_id][task_index] += 1 - else: # for env role messages - next_user_msg = msg.get("content", "") - - if new_tool_calls: - for idx, call in enumerate(new_tool_calls): - self.history[run_id][task_index].append( - {"role": "tool", "content": str(call), "tool_call_id": new_tool_call_ids[idx]}, - ) - else: - self.history[run_id][task_index].append({"role": "user", "content": next_user_msg}) - - logger.info(f"index={self.index} task_id={task_id} iteration={i}") - - if self.task_completed(run_id, task_index): - break - - reward = self.get_reward(run_id, task_index) - if self.use_memory: - if self.use_memory_addition: - new_traj_list = [ - self.get_traj_from_task_history(task_id, self.history[run_id][task_index], reward), - ] - previous_memories = self.summary_memory(new_traj_list) - if reward == 1: - self.add_memory(previous_memories) - - # update the freq & utility attributes of retrieved memories - update_utility: bool = reward == 1 - self.update_memory_information(self.retrieved_memory_list[run_id][task_index], update_utility) - - counter += 1 - if self.use_memory_deletion and counter % self.delete_freq == 0: - self.delete_memory() - - t_result = { - "run_id": run_id, - "task_id": self.task_ids[task_index], - "experiment_name": self.experiment_name, - "task_completed": self.task_completed(run_id, task_index), - "reward": reward, - "task_history": self.history[run_id][task_index], - "task_start_time": start_time, - } - if reward == 1: - break - - except Exception as e: - logger.exception(f"encounter error with {e.args}") - result.append(t_result) - return result - - def task_completed(self, run_id, index): - """ - Check if task is completed. - - Returns: - True if task is completed, False otherwise - """ - return self.history[run_id][index][-1]["content"] == "[CONVERSATION_COMPLETED]" - - -def main(): - """Main function to run the BFCLAgent.""" - with open(os.getenv("BFCL_DATA_PATH"), "r", encoding="utf-8") as f: - task_ids = [json.loads(l)["id"] for l in f] - dataset_name = "dev" - agent = BFCLAgent( - index=0, - task_ids=[task_ids[0]], - experiment_name=f"qwen3_8b_{dataset_name}", - ) - result = agent.execute() - logger.info(f"result={json.dumps(result)}") - - -if __name__ == "__main__": - main() diff --git a/benchmark/bfcl/bfcl_utils.py b/benchmark/bfcl/bfcl_utils.py deleted file mode 100644 index b20b6399..00000000 --- a/benchmark/bfcl/bfcl_utils.py +++ /dev/null @@ -1,399 +0,0 @@ -"""Utils for evaluation on BFCL tasks""" - -import json -from pathlib import Path -from typing import Dict, List, Any - -from bfcl_eval.constants.default_prompts import ( - DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC, -) -from bfcl_eval.constants.type_mappings import GORILLA_TO_OPENAPI -from bfcl_eval.eval_checker.multi_turn_eval.multi_turn_utils import ( - execute_multi_turn_func_call, -) -from bfcl_eval.model_handler.model_style import ModelStyle -from bfcl_eval.model_handler.utils import ( - convert_to_tool, - default_decode_execute_prompting, - func_doc_language_specific_pre_processing, -) - - -def load_test_case(data_path: str, test_id: str | None) -> Dict[str, Any]: - """ - load test cases by id - """ - if not Path(data_path).exists(): - raise FileNotFoundError(f"BFCL data file '{data_path}' not found") - - if test_id is None: - raise ValueError("task_id is required") - - with open(data_path, "r", encoding="utf-8") as f: - if str(test_id).isdigit(): # pylint: disable=R1720 - idx = int(test_id) - for line_no, line in enumerate(f): - if line_no == idx: - return json.loads(line) - raise ValueError(f"Test case index {idx} not found in {data_path}") - else: - for line in f: - data = json.loads(line) - if data.get("id") == test_id: - return data - raise ValueError(f"Test case id '{test_id}' not found in {data_path}") - - -def handle_user_turn( - test_entry: Dict[str, Any], - current_turn: int, -) -> Dict[str, Any]: - """ - Handle user turn by returning appropriate content from test_entry["question"]. - For non-first turns, processes user query and tools. - - Args: - test_entry: Test entry containing conversation data - current_turn: Current turn number - - Returns: - Response containing next user message and tools - """ - try: - current_turn_message = [] - tools = compile_tools(test_entry) - questions = test_entry.get("question", []) - holdout_function = test_entry.get("holdout_function", {}) - - if str(current_turn) in holdout_function: - test_entry["function"].extend(holdout_function[str(current_turn)]) - tools = compile_tools(test_entry) - assert len(questions[current_turn]) == 0, "Holdout turn should not have user message." - current_turn_message = [ - { - "role": "user", - "content": DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC, - }, - ] - return create_user_response(current_turn_message, tools) - if current_turn >= len(questions): - return create_completion_response() - - current_turn_message = questions[current_turn] - - return create_user_response(current_turn_message, tools) - - except Exception as e: - return create_error_response(f"Failed to process user message: {str(e)}") - - -def handle_tool_calls( # pylint: disable=W0613 - tool_calls: List[Dict[str, Any]], - decoded_calls: list[str], - test_entry: Dict[str, Any], - current_turn: int, -) -> Dict[str, Any]: - """ - Handle tool calls from assistant. - - Args: - tool_calls: List of tool calls in OpenAI format - decoded_calls: List of decoded function calls - test_entry: Test entry containing environment data - current_turn: Current turn number - - Returns: - Response containing tool execution results - """ - execution_results, _ = execute_multi_turn_func_call( - func_call_list=decoded_calls, - initial_config=test_entry["initial_config"], - involved_classes=test_entry["involved_classes"], - model_name="env_handler", - test_entry_id=test_entry["id"], - long_context=("long_context" in test_entry["id"] or "composite" in test_entry["id"]), - is_evaL_run=False, - ) - # print('execution_results in handler_tool_calls:', execution_results) - - return create_tool_response(tool_calls, execution_results) - - -def compile_tools(test_entry: dict) -> list: - """ - Compile functions into tools format. - - Args: - test_entry: Test entry containing functions - - Returns: - List of tools in OpenAI format - """ - functions: list = test_entry["function"] - test_category: str = test_entry["id"].rsplit("_", 1)[0] - - functions = func_doc_language_specific_pre_processing(functions, test_category) - tools = convert_to_tool(functions, GORILLA_TO_OPENAPI, ModelStyle.OpenAI_Completions) - - return tools - - -def create_tool_response( - tool_calls: List[Dict[str, Any]], - execution_results: List[str], -) -> Dict[str, Any]: - """ - Create response for tool calls. - - Args: - tool_calls: List of tool calls - execution_results: List of execution results - - Returns: - Response containing tool execution results - """ - tool_messages = [] - for i, (tool_call, result) in enumerate(zip(tool_calls, execution_results)): - tool_messages.append( - { - "role": "tool", - "content": result, - "tool_call_id": tool_call.get("id", f"call_{i}"), - }, - ) - - return {"messages": tool_messages} - - -def create_user_response( - question_turn: List[Dict[str, Any]], - tools: List[Dict[str, Any]], -) -> Dict[str, Any]: - """ - Create response containing user message. - - Args: - question_turn: List of messages for current turn - tools: List of available tools - - Returns: - Response containing user message and tools - """ - user_content = "" - for msg in question_turn: - if msg["role"] == "user": - user_content = msg["content"] - break - - return {"messages": [{"role": "user", "content": user_content}], "tools": tools} - - -def create_completion_response() -> Dict[str, Any]: - """ - Create response indicating conversation completion. - - Returns: - Response with completion message - """ - return {"messages": [{"role": "env", "content": "[CONVERSATION_COMPLETED]"}]} - - -def create_error_response(error_message: str) -> Dict[str, Any]: - """ - Create response for error conditions. - - Args: - error_message: Error message to include - - Returns: - Response containing error message - """ - return {"messages": [{"role": "env", "content": f"[ERROR] {error_message}"}]} - - -def decode_execute(result): - """ - Decode execute results for compatibility with evaluation framework. - - Args: - result: Result to decode - - Returns: - List of decoded function calls - """ - return default_decode_execute_prompting(result) - - -def extract_single_turn_response(messages: List[Dict[str, Any]]) -> str: - """ - Extract single-turn response from conversation messages. - - Args: - messages: List of conversation messages - - Returns: - String representation of the response - """ - for message in reversed(messages): - if message["role"] == "assistant": - if "tool_calls" in message and message["tool_calls"]: - formatted_calls = [] - for tool_call in message["tool_calls"]: - formatted_call = format_single_tool_call_for_eval( - tool_call, - ) - if formatted_call: - formatted_calls.append(formatted_call) - return "\n".join(formatted_calls) if formatted_calls else "" - elif message.get("content"): - return message["content"] - - return "" - - -def extract_multi_turn_responses( - messages: List[Dict[str, Any]], -) -> List[List[str]]: - """ - Extract multi-turn responses from conversation messages. - - Args: - messages: List of conversation messages - - Returns: - List of turns, each turn is a list of function call strings - """ - turns_data = [] - current_turn_responses = [] - - i = 0 - while i < len(messages): - message = messages[i] - - if message["role"] == "user": - if current_turn_responses: - turns_data.append(current_turn_responses) - current_turn_responses = [] - - i += 1 - while i < len(messages) and messages[i]["role"] == "assistant": - assistant_msg = messages[i] - - if "tool_calls" in assistant_msg and assistant_msg["tool_calls"]: - for tool_call in assistant_msg["tool_calls"]: - formatted_call = format_single_tool_call_for_eval( - tool_call, - ) - if formatted_call: - current_turn_responses.append(formatted_call) - - i += 1 - - while i < len(messages) and messages[i]["role"] == "tool": - i += 1 - else: - i += 1 - - if current_turn_responses: - turns_data.append(current_turn_responses) - - return turns_data - - -def format_single_tool_call_for_eval(tool_call: Dict[str, Any]) -> str: - """ - Format a single tool call into string representation for evaluation. - - Args: - tool_call: Single tool call in OpenAI format - - Returns: - Formatted string representation - """ - function = tool_call.get("function", {}) - function_name = function.get("name", "") - - try: - arguments = function.get("arguments", "{}") - if isinstance(arguments, str): - args_dict = json.loads(arguments) - else: - args_dict = arguments - - args_str = ", ".join([f"{k}={repr(v)}" for k, v in args_dict.items()]) - return f"{function_name}({args_str})" - - except Exception: - return f"{function_name}()" - - -def capture_and_print_score_files( - score_dir: Path, - model_name: str, - test_category: str, - eval_type: str, -): - """ - Capture and print contents of score files written to score_dir. - - Args: - score_dir: Directory containing score files - model_name: Name of the model - test_category: Category of the test - eval_type: Type of evaluation (relevance/multi_turn/single_turn) - """ - try: - print(f"\n=== {eval_type.upper()} Evaluation Result Files ===") - print(f"Model: {model_name}") - print(f"Test Category: {test_category}") - print(f"Evaluation Type: {eval_type}") - - for file_path in score_dir.rglob("*"): - if file_path.is_file(): - relative_path = file_path.relative_to(score_dir) - print(f"\n--- File: {relative_path} ---") - - try: - with open(file_path, "r", encoding="utf-8") as f: - content = f.read() - - if ( - file_path.suffix == ".json" - or content.strip().startswith("{") - or content.strip().startswith("[") - ): - try: - lines = content.strip().split("\n") - formatted_lines = [] - for line in lines: - if line.strip(): - parsed = json.loads(line) - formatted_lines.append( - json.dumps( - parsed, - ensure_ascii=False, - indent=2, - ), - ) - content = "\n".join(formatted_lines) - except json.JSONDecodeError: - pass - - print(content) - - except UnicodeDecodeError: - print(f"[Binary file, size: {file_path.stat().st_size} bytes]") - except Exception as e: - print(f"[Error reading file: {str(e)}]") - - print(f"=== {eval_type.upper()} Evaluation Result Files End ===\n") - - except Exception as e: - print(f"Error capturing evaluation result files: {str(e)}") - - -def extract_tool_schema(tools): - """Reformat tool schema""" - for i in range(len(tools)): # pylint: disable=C0200 - tools[i]["function"].pop("response") - return tools diff --git a/benchmark/bfcl/default_ids.py b/benchmark/bfcl/default_ids.py deleted file mode 100644 index 43f4065e..00000000 --- a/benchmark/bfcl/default_ids.py +++ /dev/null @@ -1,206 +0,0 @@ -# pylint: disable=C0114 -DEFAULT_TRAIN_IDS: set[str] = { - "multi_turn_base_102", - "multi_turn_base_107", - "multi_turn_base_110", - "multi_turn_base_114", - "multi_turn_base_115", - "multi_turn_base_118", - "multi_turn_base_122", - "multi_turn_base_123", - "multi_turn_base_128", - "multi_turn_base_13", - "multi_turn_base_130", - "multi_turn_base_132", - "multi_turn_base_133", - "multi_turn_base_143", - "multi_turn_base_144", - "multi_turn_base_146", - "multi_turn_base_15", - "multi_turn_base_158", - "multi_turn_base_169", - "multi_turn_base_17", - "multi_turn_base_172", - "multi_turn_base_176", - "multi_turn_base_182", - "multi_turn_base_187", - "multi_turn_base_197", - "multi_turn_base_199", - "multi_turn_base_22", - "multi_turn_base_23", - "multi_turn_base_24", - "multi_turn_base_36", - "multi_turn_base_40", - "multi_turn_base_44", - "multi_turn_base_47", - "multi_turn_base_48", - "multi_turn_base_5", - "multi_turn_base_51", - "multi_turn_base_59", - "multi_turn_base_63", - "multi_turn_base_65", - "multi_turn_base_66", - "multi_turn_base_67", - "multi_turn_base_68", - "multi_turn_base_70", - "multi_turn_base_75", - "multi_turn_base_77", - "multi_turn_base_78", - "multi_turn_base_79", - "multi_turn_base_81", - "multi_turn_base_83", - "multi_turn_base_93", -} - -DEFAULT_VAL_IDS: set[str] = { - "multi_turn_base_0", - "multi_turn_base_1", - "multi_turn_base_10", - "multi_turn_base_100", - "multi_turn_base_101", - "multi_turn_base_103", - "multi_turn_base_104", - "multi_turn_base_105", - "multi_turn_base_106", - "multi_turn_base_108", - "multi_turn_base_109", - "multi_turn_base_11", - "multi_turn_base_111", - "multi_turn_base_112", - "multi_turn_base_113", - "multi_turn_base_116", - "multi_turn_base_117", - "multi_turn_base_119", - "multi_turn_base_12", - "multi_turn_base_120", - "multi_turn_base_121", - "multi_turn_base_124", - "multi_turn_base_125", - "multi_turn_base_126", - "multi_turn_base_127", - "multi_turn_base_129", - "multi_turn_base_131", - "multi_turn_base_134", - "multi_turn_base_135", - "multi_turn_base_136", - "multi_turn_base_137", - "multi_turn_base_138", - "multi_turn_base_139", - "multi_turn_base_14", - "multi_turn_base_140", - "multi_turn_base_141", - "multi_turn_base_142", - "multi_turn_base_145", - "multi_turn_base_147", - "multi_turn_base_148", - "multi_turn_base_149", - "multi_turn_base_150", - "multi_turn_base_151", - "multi_turn_base_152", - "multi_turn_base_153", - "multi_turn_base_154", - "multi_turn_base_155", - "multi_turn_base_156", - "multi_turn_base_157", - "multi_turn_base_159", - "multi_turn_base_16", - "multi_turn_base_160", - "multi_turn_base_161", - "multi_turn_base_162", - "multi_turn_base_163", - "multi_turn_base_164", - "multi_turn_base_165", - "multi_turn_base_166", - "multi_turn_base_167", - "multi_turn_base_168", - "multi_turn_base_170", - "multi_turn_base_171", - "multi_turn_base_173", - "multi_turn_base_174", - "multi_turn_base_175", - "multi_turn_base_177", - "multi_turn_base_178", - "multi_turn_base_179", - "multi_turn_base_18", - "multi_turn_base_180", - "multi_turn_base_181", - "multi_turn_base_183", - "multi_turn_base_184", - "multi_turn_base_185", - "multi_turn_base_186", - "multi_turn_base_188", - "multi_turn_base_189", - "multi_turn_base_19", - "multi_turn_base_190", - "multi_turn_base_191", - "multi_turn_base_192", - "multi_turn_base_193", - "multi_turn_base_194", - "multi_turn_base_195", - "multi_turn_base_196", - "multi_turn_base_198", - "multi_turn_base_2", - "multi_turn_base_20", - "multi_turn_base_21", - "multi_turn_base_25", - "multi_turn_base_26", - "multi_turn_base_27", - "multi_turn_base_28", - "multi_turn_base_29", - "multi_turn_base_3", - "multi_turn_base_30", - "multi_turn_base_31", - "multi_turn_base_32", - "multi_turn_base_33", - "multi_turn_base_34", - "multi_turn_base_35", - "multi_turn_base_37", - "multi_turn_base_38", - "multi_turn_base_39", - "multi_turn_base_4", - "multi_turn_base_41", - "multi_turn_base_42", - "multi_turn_base_43", - "multi_turn_base_45", - "multi_turn_base_46", - "multi_turn_base_49", - "multi_turn_base_50", - "multi_turn_base_52", - "multi_turn_base_53", - "multi_turn_base_54", - "multi_turn_base_55", - "multi_turn_base_56", - "multi_turn_base_57", - "multi_turn_base_58", - "multi_turn_base_6", - "multi_turn_base_60", - "multi_turn_base_61", - "multi_turn_base_62", - "multi_turn_base_64", - "multi_turn_base_69", - "multi_turn_base_7", - "multi_turn_base_71", - "multi_turn_base_72", - "multi_turn_base_73", - "multi_turn_base_74", - "multi_turn_base_76", - "multi_turn_base_8", - "multi_turn_base_80", - "multi_turn_base_82", - "multi_turn_base_84", - "multi_turn_base_85", - "multi_turn_base_86", - "multi_turn_base_87", - "multi_turn_base_88", - "multi_turn_base_89", - "multi_turn_base_9", - "multi_turn_base_90", - "multi_turn_base_91", - "multi_turn_base_92", - "multi_turn_base_94", - "multi_turn_base_95", - "multi_turn_base_96", - "multi_turn_base_97", - "multi_turn_base_98", - "multi_turn_base_99", -} diff --git a/benchmark/bfcl/init_task_memory_pool.py b/benchmark/bfcl/init_task_memory_pool.py deleted file mode 100644 index 2a4fcc87..00000000 --- a/benchmark/bfcl/init_task_memory_pool.py +++ /dev/null @@ -1,235 +0,0 @@ -# pylint: disable=W0621,W1514 -"""Init task memory pool""" -import argparse -import json -from collections import defaultdict -from concurrent.futures import ThreadPoolExecutor, as_completed -from pathlib import Path -from typing import List, Dict, Any - -import requests - - -def load_task_case(data_path: str, task_id: str | None) -> Dict[str, Any]: - """ - load training cases by id - """ - if not Path(data_path).exists(): - raise FileNotFoundError(f"BFCL data file '{data_path}' not found") - - if task_id is None: - raise ValueError("task_id is required") - - with open(data_path, "r", encoding="utf-8") as f: - if str(task_id).isdigit(): # pylint: disable=R1720 - idx = int(task_id) - for line_no, line in enumerate(f): - if line_no == idx: - return json.loads(line) - raise ValueError(f"Task case index {idx} not found in {data_path}") - else: - for line in f: - data = json.loads(line) - if data.get("id") == task_id: - return data - raise ValueError(f"Task case id '{task_id}' not found in {data_path}") - - -def get_tool_prompt(tools): - """Construct prompt with provided tools""" - tool_prompt = ( - "\n\n# Tools\n\nYou may call one or more functions to assist with the user query." - "\n\nYou are provided with function signatures within XML tags:\n" - ) - for tool in tools: - tool_prompt += "\n" + json.dumps(tool) - tool_prompt += ( - "\n\n\nFor each function call, return a json object with function name" - " and arguments within XML tags:" - '\n\n{"name": , "arguments": }\n' - ) - return tool_prompt - - -def group_trajectories_by_task_id(jsonl_entries: List[Dict[str, Any]]) -> List[List[Any]]: - """ - group trajectories by task_id - - Args: - jsonl_entries: JSONL entry list - - Returns: - List[List[Any]]: trajectory list grouped by task_id - """ - grouped = defaultdict(list) - - for entry in jsonl_entries: - task_id = entry.get("task_id", "") - taks_case = load_task_case("data/multiturn_data_base.jsonl", task_id) - tools = taks_case.get("tools", [{}]) - from bfcl_utils import extract_tool_schema - - tool_schema = extract_tool_schema(tools) - entry["task_history"][0]["content"] += get_tool_prompt(tool_schema) - grouped[task_id].append(entry) - - # retain only the two with the highest and lowest rewards - filtered_groups = [] - for _, trajectories in grouped.items(): - if len(trajectories) == 1: - # when only one trajectory, retain it - filtered_groups.append(trajectories) - elif len(trajectories) == 2: - # when there are two trajectories, retain them - filtered_groups.append(trajectories) - else: - # when there are more than two trajectories, choose the two with the highest and lowest rewards - trajectories.sort(key=lambda t: t["reward"]) - min_reward_traj = trajectories[0] # highest reward - max_reward_traj = trajectories[-1] # lowest reward - filtered_groups.append([min_reward_traj, max_reward_traj]) - - return filtered_groups - - -def post_to_summarizer(trajectories: List[Any], service_url: str) -> Dict[str, Any]: - """ - post trajectories to summarizer service - - Args: - trajectories: trajectory list - service_url: summarizer service URL - - Returns: - response json - """ - trajectory_dicts = [ - { - "task_id": traj["task_id"], - "messages": traj["task_history"], - "score": traj["reward"], - } - for traj in trajectories - ] - - request_data = { - "trajectories": trajectory_dicts, - "success_threshold": 1.0, - "enable_soft_comparison": True, - "validation_threshold": 0.5, - } - - try: - response = requests.post(f"{service_url}/summary_task_memory", json=request_data) - response.raise_for_status() - return response.json() - except Exception as e: - return {"error": str(e), "trajectories_count": len(trajectories)} - - -def process_trajectories_with_threads( - grouped_trajectories: List[List[Any]], - service_url: str, - n_threads: int = 4, -) -> List[Dict[str, Any]]: - """ - use threads to process trajectories - - Args: - grouped_trajectories: group trajectory list by task_id - service_url: memory summarizer service URL - n_threads: number of threads - - Returns: - all results - """ - results = [] - - with ThreadPoolExecutor(max_workers=n_threads) as executor: - future_to_group = { - executor.submit(post_to_summarizer, group, service_url): i for i, group in enumerate(grouped_trajectories) - } - - for future in as_completed(future_to_group): - group_index = future_to_group[future] - try: - result = future.result() - result["group_index"] = group_index - result["group_size"] = len(grouped_trajectories[group_index]) - results.append(result) - if "memory_list" in result["metadata"]: - print(f'✅ Group {group_index} processed: {result["metadata"].get("memory_list", 0)}') - memory_list = result["metadata"].get("memory_list", []) - response = requests.post(url=f"{service_url}/add_task_memory", json={"memory_list": memory_list}) - response.raise_for_status() - else: - print(f"❌ Group {group_index} processed: error") - except Exception as e: - error_result = { - "group_index": group_index, - "group_size": len(grouped_trajectories[group_index]), - "error": str(e), - } - results.append(error_result) - print(f"❌ Group {group_index} failed: {e}") - - return results - - -def main(): - """Main function to convert JSONL to memories using ReMe service.""" - parser = argparse.ArgumentParser(description="Convert JSONL to memories using ReMe service") - parser.add_argument("--jsonl_file", type=str, required=True, help="Path to the JSONL file") - parser.add_argument("--service_url", type=str, default="http://localhost:8002", help="ReMe service URL") - parser.add_argument("--output_file", type=str, help="Output file to save results (optional)") - parser.add_argument("--n_threads", type=int, default=4, help="Number of threads for processing") - - args = parser.parse_args() - - print(f"Processing JSONL file: {args.jsonl_file}") - print(f"Service URL: {args.service_url}") - print(f"Threads: {args.n_threads}") - - with open(args.jsonl_file, "r") as f: - data = [json.loads(line) for line in f] - print(f"Loaded {len(data)} entries from JSONL file") - - grouped_trajectories = group_trajectories_by_task_id(data) - print(f"Total groups: {len(grouped_trajectories)}") - - results = process_trajectories_with_threads( - grouped_trajectories, - args.service_url, - n_threads=args.n_threads, - ) - - print(f"Processed {len(results)} groups") - - success_count = sum(1 for r in results if "error" not in r) - error_count = len(results) - success_count - total_memories = sum(len(r["metadata"].get("memory_list", [])) for r in results if "memory_list" in r["metadata"]) - - print(f"✅ Success: {success_count}") - print(f"❌ Errors: {error_count}") - print(f"📊 Total task memories created: {total_memories}") - - if args.output_file: - try: - summary = { - "jsonl_file": args.jsonl_file, - "total_groups": len(grouped_trajectories), - "success_count": success_count, - "error_count": error_count, - "total_task_memories": total_memories, - "results": results, - } - - with open(args.output_file, "w") as f: - json.dump(summary, f, indent=2) - print(f"Results saved to: {args.output_file}") - except Exception as e: - print(f"Error saving results: {e}") - - -if __name__ == "__main__": - main() diff --git a/benchmark/bfcl/preprocess.py b/benchmark/bfcl/preprocess.py deleted file mode 100644 index e95a9d42..00000000 --- a/benchmark/bfcl/preprocess.py +++ /dev/null @@ -1,73 +0,0 @@ -# pylint: disable=W0621 -"""Preprocess multi-turn test cases""" - -import json - - -from pathlib import Path -from bfcl_eval.model_handler.model_style import ModelStyle -from bfcl_eval.eval_checker.eval_runner_helper import load_file -from bfcl_eval.constants.type_mappings import GORILLA_TO_OPENAPI -from bfcl_eval.constants.eval_config import MULTI_TURN_FUNC_DOC_PATH -from bfcl_eval.constants.category_mapping import MULTI_TURN_FUNC_DOC_FILE_MAPPING -from bfcl_eval.model_handler.utils import ( - convert_to_tool, - func_doc_language_specific_pre_processing, -) - - -def process_multi_turn_test_case(file_path, output_path): - """ - Multi-turn test cases don't have the function doc in the prompt. We need to add them here. - """ - test_cases = [] - with open(output_path, "w", encoding="utf-8") as outf: - with open(file_path, encoding="utf-8") as f: - file = f.readlines() - for line in file: - entry = json.loads(line) - if "multi_turn" not in entry["id"]: - continue - test_category: str = entry["id"].rsplit("_", 1)[0] - involved_classes = entry["involved_classes"] - entry["function"] = [] - for func_collection in involved_classes: - # func_doc is a list of dict - func_doc = load_file( - MULTI_TURN_FUNC_DOC_PATH / MULTI_TURN_FUNC_DOC_FILE_MAPPING[func_collection], - ) - entry["function"].extend(func_doc) - - # Handle Miss Func category; we need to remove the holdout function doc - if "missed_function" in entry: - for turn_index, missed_func_names in entry["missed_function"].items(): - entry["missed_function"][turn_index] = [] - for missed_func_name in missed_func_names: - for i, func_doc in enumerate(entry["function"]): - if func_doc["name"] == missed_func_name: - # Add the missed function doc to the missed_function list - entry["missed_function"][turn_index].append(func_doc) - # Remove it from the function list - entry["function"].pop(i) - break - - functions = func_doc_language_specific_pre_processing(entry["function"], test_category) - tools = convert_to_tool(functions, GORILLA_TO_OPENAPI, ModelStyle.OpenAI_Completions) - - test_cases.append( - { - "id": entry["id"], - "messages": entry["question"][0], - "tools": tools, - "extra": entry, - }, - ) - outf.write(json.dumps(test_cases[-1], ensure_ascii=False) + "\n") - - return test_cases - - -if __name__ == "__main__": - file_path = Path("./gorilla/berkeley-function-call-leaderboard/bfcl_eval/data/BFCL_v3_multi_turn_base.json") - output_path = "data/multiturn_data_base.jsonl" - preprocessed_test_cases = process_multi_turn_test_case(file_path, output_path) diff --git a/benchmark/bfcl/run_bfcl.py b/benchmark/bfcl/run_bfcl.py deleted file mode 100644 index 6ea8c325..00000000 --- a/benchmark/bfcl/run_bfcl.py +++ /dev/null @@ -1,151 +0,0 @@ -"""Run evaluation on BFCL-V3-Multi-Turn-Base dataset.""" - -import time -import json -from pathlib import Path - -import ray -import requests -from loguru import logger -from dotenv import load_dotenv -from bfcl_agent import BFCLAgent - -load_dotenv("../../.env") - - -def run_agent( - max_workers: int, - dataset_name: str, - experiment_suffix: str, - model_name: str = "qwen3-8b", - enable_thinking: bool = False, - data_path: str = "data/multiturn_data_base_val.jsonl", - answer_path: Path = Path("data/possible_answer"), - num_trials: int = 1, - use_memory: bool = False, - memory_base_url: str = "http://0.0.0.0:8002/", - use_memory_addition: bool = True, - use_memory_deletion: bool = False, - delete_freq: int = 10, - freq_threshold: int = 5, - utility_threshold: float = 0.5, -): - """Run the agent""" - experiment_name = dataset_name + "_" + experiment_suffix - path: Path = Path( - f"./exp_result/{model_name}/with_think" if enable_thinking else f"./exp_result/{model_name}/no_think", - ) - path.mkdir(parents=True, exist_ok=True) - - with open(data_path, "r", encoding="utf-8") as f: - task_ids = [json.loads(line)["id"] for line in f] - - result: list = [] - - def dump_file(): - with open(path / f"{experiment_name}.jsonl", "a", encoding="utf-8") as f: - for x in result: - f.write(json.dumps(x) + "\n") - - future_list: list = [] - for i in range(max_workers): - actor = BFCLAgent.remote( - index=i, - model_name=model_name, - task_ids=task_ids[i::max_workers], - experiment_name=experiment_name, - data_path=data_path, - answer_path=answer_path, - num_trials=num_trials, - use_memory=use_memory, - memory_base_url=memory_base_url, - use_memory_addition=use_memory_addition, - use_memory_deletion=use_memory_deletion, - delete_freq=delete_freq, - freq_threshold=freq_threshold, - utility_threshold=utility_threshold, - enable_thinking=enable_thinking, - ) - future = actor.execute.remote() - future_list.append(future) - time.sleep(1) - logger.info("submit complete") - - for i, future in enumerate(future_list): - t_result = ray.get(future) - if t_result: - if isinstance(t_result, list): - result.extend(t_result) - else: - result.append(t_result) - - logger.info(f"{i + 1}/{len(task_ids)} complete") - dump_file() - - -def handle_api_response(response: requests.Response): - """Handle API response with proper error checking""" - if response.status_code != 200: - print(f"Error: {response.status_code}") - print(response.text) - return None - - return response.json() - - -def load_memory(path: str = "docs/library", api_url: str = "http://0.0.0.0:8002/"): - """Load memories from disk into the vector store""" - response = requests.post( - url=f"{api_url}load_memory", - json={ - "load_file_path": path, - "clear_existing": True, - }, - ) - - result = handle_api_response(response) - if result: - print(f"Memory loaded from {path}") - - -def main(): - """Main function""" - max_workers = 4 - if max_workers > 1: - ray.init(num_cpus=max_workers) - - num_runs = 4 - num_trials = 1 - model_name = "qwen3-8b" - enable_thinking = True - use_memory = True - use_memory_addition = False - use_memory_deletion = False - memory_base_url = "http://0.0.0.0:8003/" - - if use_memory: - load_file_path = "docs/library/paper_data/task/bfcl_qwen3_8b.jsonl" - load_memory(load_file_path, memory_base_url) - - for _ in range(num_runs): - run_agent( - max_workers=max_workers, - model_name=model_name, - dataset_name="bfcl-multi-turn-base-val", - experiment_suffix="w-fixed-memory", - data_path="data/multiturn_data_base_val.jsonl", - answer_path=Path("data/possible_answer"), - enable_thinking=enable_thinking, - num_trials=num_trials, - use_memory=use_memory, - memory_base_url=memory_base_url, - use_memory_addition=use_memory_addition, - use_memory_deletion=use_memory_deletion, - delete_freq=5, - freq_threshold=5, - utility_threshold=0.5, - ) - - -if __name__ == "__main__": - main() diff --git a/benchmark/bfcl/run_exp_statistic.py b/benchmark/bfcl/run_exp_statistic.py deleted file mode 100644 index 18efcc8d..00000000 --- a/benchmark/bfcl/run_exp_statistic.py +++ /dev/null @@ -1,163 +0,0 @@ -"""Run the experiment statistic.""" - -import json -from collections import defaultdict -from pathlib import Path - -import pandas as pd -from loguru import logger - - -def calculate_best_at_k(scores: list, k: int) -> float: - """ - Calculate best@k - Divide scores into groups of size k, take the maximum value in each group, - then average these maximum values - - Args: - scores: List of after_score values for all runs of a task - k: Group size - - Returns: - best@k value - """ - if len(scores) % k != 0: - raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})") - - group_maxs = [] - for i in range(0, len(scores), k): - group = scores[i : i + k] - group_maxs.append(max(group)) - - return sum(group_maxs) / len(group_maxs) - - -def calculate_pass_at_k(scores: list, k: int) -> float: - """Calculate pass@k.""" - if len(scores) % k != 0: - raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})") - - group_maxs = [] - for i in range(0, len(scores), k): - group = scores[i : i + k] - is_pass = 1.0 if max(group) >= 1.0 else 0.0 - group_maxs.append(is_pass) - - return sum(group_maxs) / len(group_maxs) - - -def get_possible_k_values(total_runs: int) -> list: - """ - Get all possible k values (factors of total_runs) - - Args: - total_runs: Total number of runs - - Returns: - List of k values in descending order - """ - k_values = [] - for k in range(1, total_runs + 1): - if total_runs % k == 0: - k_values.append(k) - return sorted(k_values, reverse=True) # Sort from large to small - - -def run_exp_statistic(): - """Run the experiment statistic.""" - path: Path = Path("./exp_result/qwen3-8b/with_think") - - # Store results for all experiments - all_results = {} - for file in path.glob("*.jsonl"): - # Group results by task_id - task_results = defaultdict(list) - print(file) - with open(file, "r", encoding="utf-8") as f: - for line in f: - if not line.strip(): - continue - data = json.loads(line) - - if isinstance(data, list): - for part_data in data: - task_id = part_data["task_id"] - after_score = part_data["reward"] - task_results[task_id].append(after_score) - else: - task_id = data["task_id"] - after_score = data["reward"] - task_results[task_id].append(after_score) - - if not task_results: - logger.warning(f"No valid data found in file {file}") - continue - - # Check if each task has consistent number of runs - run_counts = [len(scores) for scores in task_results.values()] - if len(set(run_counts)) > 1: - logger.warning(f"Inconsistent number of runs for different tasks in file {file}: {set(run_counts)}") - continue - - num_runs = run_counts[0] - logger.info(f"File {file}: {len(task_results)} tasks, {num_runs} runs per task") - - # Get all possible k values - k_values = get_possible_k_values(num_runs) - logger.info(f"Calculable best@k values: {k_values}") - - # Calculate various best@k values - file_results = {"file": file.name} - - for k in k_values: - best_at_k_scores = [] - pass_at_k_scores = [] - for task_id, scores in task_results.items(): - try: - best_k_score = calculate_best_at_k(scores, k) - pass_at_k_score = calculate_pass_at_k(scores, k) - pass_at_k_scores.append(pass_at_k_score) - best_at_k_scores.append(best_k_score) - except ValueError as e: - logger.error(f"Error calculating best@{k} for task {task_id}: {e}") - continue - - if best_at_k_scores: - avg_best_at_k = sum(best_at_k_scores) / len(best_at_k_scores) - file_results[f"best@{k}"] = avg_best_at_k - logger.info(f"file={file.name} best@{k}={avg_best_at_k:.4f}") - - if pass_at_k_scores: - avg_pass_at_k = sum(pass_at_k_scores) / len(pass_at_k_scores) - file_results[f"pass@{k}"] = avg_pass_at_k - logger.info(f"file={file.name} pass@{k}={avg_pass_at_k:.4f}") - - all_results[file.name] = file_results - - # Create and display table - if all_results: - df = pd.DataFrame(list(all_results.values())) - df = df.set_index("file") - - # Sort columns by the number in column name (best@8, best@4, best@2, best@1) - # best_columns = [col for col in df.columns if col.startswith('best@')] - best_columns = list(df.columns) - best_columns.sort(key=lambda x: x, reverse=False) - df = df[best_columns] - - print("\n" + "=" * 80) - print("Experiment Results Summary Table") - print("=" * 80) - print(df.round(4)) - print("=" * 80) - - # Save table to CSV - output_path = path / "experiment_summary.csv" - df.to_csv(output_path) - logger.info(f"Results table saved to: {output_path}") - else: - logger.warning("No valid experiment results found") - - -if __name__ == "__main__": - run_exp_statistic() diff --git a/benchmark/bfcl/split_into_trainval.py b/benchmark/bfcl/split_into_trainval.py deleted file mode 100644 index e217def7..00000000 --- a/benchmark/bfcl/split_into_trainval.py +++ /dev/null @@ -1,69 +0,0 @@ -"""Split the JSONL file into train and validation sets.""" - -import argparse -import json -import random - -from default_ids import DEFAULT_TRAIN_IDS, DEFAULT_VAL_IDS - - -def split_jsonl( - input_file: str, - train_file: str, - val_file: str, - ratio: float = 0.75, - random_split: bool = False, -) -> None: - """Split the JSONL file into train and validation sets.""" - with open(input_file, "r", encoding="utf-8") as f: - data = [json.loads(line) for line in f] - - if random_split: - random.shuffle(data) - split_idx = int(len(data) * ratio) - train_data = data[:split_idx] - val_data = data[split_idx:] - else: - train_data = [] - val_data = [] - unknown_ids: list[str] = [] - for obj in data: - if "id" not in obj: - raise ValueError(f"Missing 'id' field in input file: {input_file}") - obj_id = str(obj["id"]) - if obj_id in DEFAULT_TRAIN_IDS: - train_data.append(obj) - elif obj_id in DEFAULT_VAL_IDS: - val_data.append(obj) - else: - unknown_ids.append(obj_id) - - if len(train_data) + len(val_data) != len(data): - missing = len(data) - (len(train_data) + len(val_data)) - examples = ", ".join(unknown_ids) if unknown_ids else "(none)" - raise ValueError( - f"{missing} samples in {input_file} not found in train_ref/val_ref id sets. Examples: {examples}", - ) - - with open(train_file, "w", encoding="utf-8") as f: - for item in train_data: - f.write(json.dumps(item, ensure_ascii=False) + "\n") - with open(val_file, "w", encoding="utf-8") as f: - for item in val_data: - f.write(json.dumps(item, ensure_ascii=False) + "\n") - - -if __name__ == "__main__": - parser = argparse.ArgumentParser(description="Split JSONL file into train and validation sets.") - parser.add_argument("--input", required=True, help="Path to input JSONL file") - parser.add_argument("--train", required=True, help="Path to output train file") - parser.add_argument("--val", required=True, help="Path to output validation file") - parser.add_argument("--ratio", type=float, default=0.5, help="Train ratio (default: 0.8)") - parser.add_argument( - "--random", - action="store_true", - help="Whether to randomly split input into train/val. " - "If false, split strictly by default train/val id sets (see default_ids.py).", - ) - args = parser.parse_args() - split_jsonl(args.input, args.train, args.val, args.ratio, args.random) diff --git a/benchmark/halumem/eval_reme.py b/benchmark/halumem/eval_reme.py deleted file mode 100644 index 8cc27e53..00000000 --- a/benchmark/halumem/eval_reme.py +++ /dev/null @@ -1,1494 +0,0 @@ -""" -HaluMem Benchmark Evaluator for ReMe - Question Answering - -A modular evaluation pipeline that: -1. Loads HaluMem benchmark data -2. Processes user sessions through ReMe (summarization + retrieval) -3. Evaluates question answering performance -4. Generates comprehensive metrics - -Usage: - python benchmark/halumem/eval_reme.py \ - --data_path /Users/yuli/workspace/HaluMem/data/HaluMem-Medium.jsonl \ - --top_k 20 --user_num 100 --max_concurrency 20 -""" - -import asyncio -import json -import os -import re -import shutil -import time -from dataclasses import dataclass -from datetime import datetime, timezone -from pathlib import Path -from typing import Any - -from loguru import logger - -from reme.reme import ReMe - - -# ==================== Configuration ==================== - - -@dataclass -class EvalConfig: - """Evaluation configuration parameters.""" - - data_path: str - top_k: int = 20 - user_num: int = 1 - max_concurrency: int = 1 - batch_size: int = 40 - output_dir: str = "bench_results/reme" - reme_model_name: str = "qwen-flash" - eval_model_name: str = "qwen3-max" - algo_version: str = "v1" - enable_thinking_params: bool = False - - -# ==================== Utilities ==================== - - -class DataLoader: - """Handles loading and parsing of HaluMem data.""" - - @staticmethod - def load_jsonl(file_path: str) -> list[dict]: - """Load all entries from a JSONL file.""" - with open(file_path, "r", encoding="utf-8") as f: - return [json.loads(line.strip()) for line in f if line.strip()] - - @staticmethod - def extract_user_name(persona_info: str) -> str: - """Extract user name from persona info string.""" - match = re.search(r"Name:\s*(.*?); Gender:", persona_info) - if not match: - raise ValueError(f"No name found in persona_info: {persona_info}") - return match.group(1).strip() - - @staticmethod - def format_dialogue_messages(dialogue: list[dict]) -> list[dict]: - """Format dialogue into ReMe message format with conversation_time (user messages only).""" - return [ - { - "role": turn["role"], - "content": turn["content"], - "time_created": datetime.strptime( - turn["timestamp"], - "%b %d, %Y, %H:%M:%S", - ) - .replace(tzinfo=timezone.utc) - .strftime("%Y-%m-%d %H:%M:%S"), - } - for turn in dialogue - if turn["role"] == "user" # Only include user messages - ] - - @staticmethod - def format_dialogue_for_eval(dialogue: list[dict], user_name: str = None) -> str: - """Format dialogue into string for evaluation.""" - formatted_turns = [] - for turn in dialogue: - timestamp = ( - datetime.strptime( - turn["timestamp"], - "%b %d, %Y, %H:%M:%S", - ) - .replace(tzinfo=timezone.utc) - .strftime("%Y-%m-%d %H:%M:%S") - ) - - # Use user_name if role is 'user' and user_name is provided - role = user_name if turn["role"] == "user" and user_name else turn["role"] - - formatted_turns.append( - f"Role: {role}\n" f"Content: {turn['content']}\n" f"Time: {timestamp}", - ) - return "\n\n".join(formatted_turns) - - -class FileManager: - """Manages file I/O operations.""" - - def __init__(self, base_dir: str): - self.base_dir = Path(base_dir) - self.tmp_dir = self.base_dir - self.tmp_dir.mkdir(parents=True, exist_ok=True) - - def get_user_dir(self, user_name: str) -> Path: - """Get the directory path for a user.""" - user_dir = self.tmp_dir / user_name - user_dir.mkdir(parents=True, exist_ok=True) - return user_dir - - def get_session_file(self, user_name: str, session_id: int) -> Path: - """Get the file path for a specific session.""" - return self.get_user_dir(user_name) / f"session_{session_id}.json" - - def save_session(self, user_name: str, session_id: int, data: dict): - """Save session data to file.""" - file_path = self.get_session_file(user_name, session_id) - with open(file_path, "w", encoding="utf-8") as f: - json.dump(data, f, ensure_ascii=False, indent=2) - logger.info(f"✅ Saved session {session_id} to {file_path}") - - def load_session(self, user_name: str, session_id: int) -> dict | None: - """Load session data from file.""" - file_path = self.get_session_file(user_name, session_id) - if not file_path.exists(): - return None - with open(file_path, "r", encoding="utf-8") as f: - return json.load(f) - - def user_has_cache(self, user_name: str) -> bool: - """Check if user has cached results.""" - user_dir = self.get_user_dir(user_name) - return any(f.name.startswith("session_") and f.suffix == ".json" for f in user_dir.iterdir()) - - def combine_results(self, output_file: str): - """Combine all user session files into a single JSONL file.""" - with open(output_file, "w", encoding="utf-8") as f_out: - for user_dir in self.tmp_dir.iterdir(): - if not user_dir.is_dir(): - continue - - session_files = sorted( - [f for f in user_dir.iterdir() if f.name.startswith("session_") and f.suffix == ".json"], - ) - - if not session_files: - continue - - # Load first session to get user metadata - with open(session_files[0], "r", encoding="utf-8") as f_in: - first_session = json.load(f_in) - - user_data = { - "uuid": first_session["uuid"], - "user_name": first_session["user_name"], - "sessions": [], - } - - # Load all sessions - for session_file in session_files: - with open(session_file, "r", encoding="utf-8") as f_in: - session_data = json.load(f_in) - # Remove redundant user metadata - session_data.pop("uuid", None) - session_data.pop("user_name", None) - user_data["sessions"].append(session_data) - - f_out.write(json.dumps(user_data, ensure_ascii=False) + "\n") - - -# ==================== Evaluation Functions ==================== - - -async def answer_question_with_memories( - reme: ReMe, - question: str, - memories: str, - user_id: str = None, - eval_model_name: str = "qwen3-30b-a3b-instruct-2507", -): - """ - Answer a question using retrieved memories with PROMPT_MEMZERO_JSON template. - - Args: - reme: ReMe instance with default_llm and prompt_handler - question: The question to answer - memories: The retrieved memories (formatted as context) - user_id: Optional user ID for context formatting - eval_model_name: Model name to use for LLM request - - Returns: - dict with 'reasoning' and 'answer' fields - """ - # Format context with memories - if user_id: - context = reme.prompt_handler.prompt_format( - "TEMPLATE_MEMOS", - user_id=user_id, - memories=memories, - ) - else: - context = f"Memories:\n{memories}" - - # Use PROMPT_MEMZERO_JSON template for structured JSON response - prompt = reme.prompt_handler.prompt_format( - "PROMPT_MEMZERO_JSON", - context=context, - question=question, - ) - - result = await reme.get_llm(eval_model_name).simple_request_for_json( - prompt=prompt, - model_name=None, - ) - - return result - - -async def evaluation_for_memory_accuracy( - reme: ReMe, - dialogue: str, - golden_memories: list[dict], - candidate_memory: dict, - eval_model_name: str = "qwen-flash", -): - """ - Memory Accuracy Evaluation - Check if an extracted memory is accurate. - - Args: - reme: ReMe instance with default_llm and prompt_handler - dialogue: The formatted dialogue string - golden_memories: List of golden memory points from the session - candidate_memory: The extracted memory to evaluate - eval_model_name: Model name to use for LLM request - - Returns: - dict with 'accuracy_score' (0/1/2), 'is_included_in_golden_memories' (true/false), and 'reason' - """ - # Format golden memories as string - golden_memories_text = "\n".join( - [f"- {m.get('memory_content', str(m))}" for m in golden_memories], - ) - - # Extract candidate memory content - candidate_content = candidate_memory.get("content", candidate_memory.get("memory_content", str(candidate_memory))) - - prompt = reme.prompt_handler.prompt_format( - "EVALUATION_PROMPT_FOR_MEMORY_ACCURACY", - dialogue=dialogue, - golden_memories=golden_memories_text, - candidate_memory=candidate_content, - ) - - result = await reme.get_llm(eval_model_name).simple_request_for_json( - prompt=prompt, - model_name=None, - ) - - return result - - -async def evaluation_for_memory_integrity( - reme: ReMe, - extracted_memories: list[dict], - expected_memory_point: dict, -): - """ - Memory Integrity Evaluation - Check if extracted memories cover the expected memory point. - - Args: - reme: ReMe instance with default_llm and prompt_handler - extracted_memories: List of extracted memory dicts - expected_memory_point: The expected memory point dict with 'memory_content' field - model_name: Model name to use for LLM request - - Returns: - dict with 'reasoning' and 'score' fields (score: 0, 1, or 2) - """ - # Format extracted memories as string - memories_text = "\n".join( - [f"- {m.get('content', m.get('memory_content', str(m)))}" for m in extracted_memories], - ) - - # Extract expected memory point content - expected_content = expected_memory_point.get("memory_content", str(expected_memory_point)) - - prompt = reme.prompt_handler.prompt_format( - "EVALUATION_PROMPT_FOR_MEMORY_INTEGRITY", - memories=memories_text, - expected_memory_point=expected_content, - ) - - result = await reme.get_llm("qwen-flash").simple_request_for_json( - prompt=prompt, - model_name=None, - ) - - return result - - -async def evaluation_for_question( - reme: ReMe, - question: str, - reference_answer: str, - key_memory_points: str, - response: str, - dialogue: str = None, - model_name: str = None, -): - """ - Question-Answering Evaluation with optional Dialogue Context. - - Args: - reme: ReMe instance with default_llm and prompt_handler - question: The question string to be evaluated. - reference_answer: The reference (gold-standard) answer. - key_memory_points: The memory points used to derive the reference answer. - response: The answer produced by the memory system. - dialogue: Optional formatted dialogue history (role, content, time_created). - model_name: Model name to use for LLM request - - Returns: - dict with 'reasoning' and 'evaluation_result' fields - """ - prompt = reme.prompt_handler.prompt_format( - "EVALUATION_PROMPT_FOR_QUESTION2", - question=question, - reference_answer=reference_answer, - key_memory_points=key_memory_points, - response=response, - dialogue=dialogue if dialogue else "", - ) - - result = await reme.get_llm("qwen3_max_instruct").simple_request_for_json( - prompt=prompt, - model_name=model_name, - ) - - return result - - -# ==================== Memory Operations ==================== - - -class MemoryProcessor: - """Handles ReMe memory operations.""" - - def __init__( - self, - reme: ReMe, - reme_model_name: str = "qwen3-max", - eval_model_name: str = "qwen3-max", - algo_version: str = "halumem", - enable_thinking_params: bool = False, - ): - self.reme = reme - self.reme_model_name = reme_model_name - self.eval_model_name = eval_model_name - self.algo_version = algo_version - self.enable_thinking_params = enable_thinking_params - - async def add_memories( - self, - user_id: str, - messages: list[dict], - batch_size: int = 10000, - ) -> tuple[list[str], list, float]: - """ - Add memories in batches using ReMe and return extracted memory contents. - - Returns: - tuple: (extracted_memories, agent_messages, total_duration_ms) - """ - extracted_memories = [] - summary_messages = [] - total_duration_ms = 0 - - for i in range(0, len(messages), batch_size): - batch = messages[i : i + batch_size] - start = time.time() - - # Use new summary API - result = await self.reme.summarize_memory( - messages=batch, - user_name=user_id, - version=self.algo_version, - return_dict=True, - enable_time_filter=True, - enable_thinking_params=self.enable_thinking_params, - ) - - duration_ms = (time.time() - start) * 1000 - total_duration_ms += duration_ms - - extracted_memories.extend( - [ - memory_node.model_dump(exclude_none=True) - for memory_node in result["answer"] - if "time_int" in memory_node.metadata and memory_node.when_to_use == "" - ], - ) - summary_messages.extend([m.simple_dump(enable_argument_dict=True) for m in result["messages"]]) - - return extracted_memories, summary_messages, total_duration_ms - - async def search_memory( - self, - query: str, - user_id: str, - top_k: int = 20, - ) -> tuple[dict, list, float]: - """ - Search memory using ReMe and return structured answer with reasoning. - - Returns: - tuple: (answer_dict, agent_messages, duration_ms) - answer_dict contains: {"reasoning": str, "answer": str, "memories": str} - """ - start = time.time() - - # Retrieve memories from ReMe using new API - result = await self.reme.retrieve_memory( - query=query, - retrieve_top_k=top_k, - user_name=user_id, - version=self.algo_version, - return_dict=True, - enable_time_filter=True, - enable_thinking_params=self.enable_thinking_params, - ) - - # Extract memories from response - memories = result["answer"] - agent_messages = [x.simple_dump(enable_argument_dict=True) for x in result["messages"]] - retrieved_nodes = [x.model_dump(exclude_none=True) for x in result["retrieved_nodes"]] - - # Use LLM to generate structured answer from memories - answer_result = await answer_question_with_memories( - reme=self.reme, - question=query, - memories=memories, - user_id=user_id, - eval_model_name=self.eval_model_name, - ) - - # Add original memories to the result - answer_result["memories"] = memories - answer_result["retrieved_nodes"] = retrieved_nodes - - duration_ms = (time.time() - start) * 1000 - return answer_result, agent_messages, duration_ms - - -# ==================== Evaluation ==================== - - -class QuestionAnsweringEvaluator: - """Evaluates question answering performance.""" - - def __init__(self, memory_processor: MemoryProcessor, reme: ReMe, top_k: int, eval_model_name: str = "qwen3-max"): - self.memory_processor = memory_processor - self.reme = reme - self.top_k = top_k - self.eval_model_name = eval_model_name - - async def evaluate_questions( - self, - questions: list[dict], - user_name: str, - uuid: str, - session_id: int, - formatted_dialogue: str, - ) -> list[dict]: - """Evaluate all questions for a session.""" - results = [] - - for qa in questions: - answer_dict, agent_messages, duration_ms = await self.memory_processor.search_memory( - query=qa["question"], - user_id=user_name, - top_k=self.top_k, - ) - - # Extract answer and reasoning from the structured response - system_answer = answer_dict.get("answer", "") - system_reasoning = answer_dict.get("reasoning", "") - retrieved_memories = answer_dict.get("memories", "") - retrieved_nodes = answer_dict.get("retrieved_nodes", "") - - # Evaluate response - evidence_text = "\n".join([e["memory_content"] for e in qa["evidence"]]) - eval_result = await evaluation_for_question( - reme=self.reme, - question=qa["question"], - reference_answer=qa["answer"], - key_memory_points=evidence_text, - response=system_answer, - dialogue=formatted_dialogue, - model_name=self.eval_model_name, - ) - - eval_result_original_answer = await evaluation_for_question( - reme=self.reme, - question=qa["question"], - reference_answer=qa["answer"], - key_memory_points=evidence_text, - response=retrieved_memories, - dialogue=formatted_dialogue, - model_name=self.eval_model_name, - ) - - # Build result record - qa_result = { - **qa, - "uuid": uuid, - "session_id": session_id, - "system_response": system_answer, - "system_reasoning": system_reasoning, - "retrieved_memories": retrieved_memories, - "retrieved_nodes": retrieved_nodes, - "retrieve_messages": agent_messages, - "search_duration_ms": duration_ms, - "result_type": eval_result.get("evaluation_result"), - "question_answering_reasoning": eval_result.get("reasoning", ""), - "original_result_type": eval_result_original_answer.get("evaluation_result"), - "original_question_answering_reasoning": eval_result_original_answer.get("reasoning", ""), - } - results.append(qa_result) - - return results - - -class MemoryIntegrityEvaluator: - """Evaluates memory integrity - whether extracted memories cover expected memory points.""" - - def __init__(self, reme: ReMe, eval_model_name: str = "gpt-4o-mini"): - self.reme = reme - self.eval_model_name = eval_model_name - - async def evaluate_memory_points( - self, - extracted_memories: list[dict], - memory_points: list[dict], - ) -> list[dict]: - """ - Evaluate whether extracted memories cover each expected memory point. - - Args: - extracted_memories: List of memories extracted by the system - memory_points: List of expected memory points from the session - - Returns: - List of evaluation results, one per memory point - """ - results = [] - - for memory_point in memory_points: - eval_result = await evaluation_for_memory_integrity( - reme=self.reme, - extracted_memories=extracted_memories, - expected_memory_point=memory_point, - ) - - # Build result record - if eval_result is None: - eval_result = {} - integrity_result = { - **memory_point, - "integrity_score": eval_result.get("score"), - "integrity_reasoning": eval_result.get("reasoning", ""), - } - results.append(integrity_result) - - return results - - -class MemoryAccuracyEvaluator: - """Evaluates memory accuracy - whether each extracted memory is accurate.""" - - def __init__(self, reme: ReMe, eval_model_name: str = "gpt-4o-mini"): - self.reme = reme - self.eval_model_name = eval_model_name - - async def evaluate_extracted_memories( - self, - extracted_memories: list[dict], - memory_points: list[dict], - formatted_dialogue: str, - ) -> list[dict]: - """ - Evaluate the accuracy of each extracted memory. - - Args: - extracted_memories: List of memories extracted by the system - memory_points: List of golden memory points from the session - formatted_dialogue: The formatted dialogue string - - Returns: - List of evaluation results, one per extracted memory - """ - results = [] - - for memory in extracted_memories: - eval_result = await evaluation_for_memory_accuracy( - reme=self.reme, - dialogue=formatted_dialogue, - golden_memories=memory_points, - candidate_memory=memory, - eval_model_name=self.eval_model_name, - ) - - # Build result record - if eval_result is None: - eval_result = {} - accuracy_result = { - "memory_content": memory.get("content", memory.get("memory_content", str(memory))), - "memory_id": memory.get("memory_id", ""), - "accuracy_score": eval_result.get("accuracy_score"), - "is_included_in_golden_memories": eval_result.get("is_included_in_golden_memories"), - "accuracy_reason": eval_result.get("reason", ""), - } - results.append(accuracy_result) - - return results - - -class MetricsAggregator: - """Aggregates evaluation metrics.""" - - @staticmethod - def _compute_single_metric(qa_records: list[dict], result_key: str) -> dict[str, Any]: - """Compute metrics for a single result type key.""" - total = len(qa_records) - if total == 0: - return { - "correct_qa_ratio(all)": 0, - "hallucination_qa_ratio(all)": 0, - "omission_qa_ratio(all)": 0, - "correct_qa_ratio(valid)": 0, - "hallucination_qa_ratio(valid)": 0, - "omission_qa_ratio(valid)": 0, - "qa_valid_num": 0, - "qa_num": 0, - } - - correct = 0 - hallucination = 0 - omission = 0 - valid = 0 - - for qa in qa_records: - result_type = qa.get(result_key, "") - - if result_type in ["Correct", "Hallucination", "Omission"]: - valid += 1 - if result_type == "Correct": - correct += 1 - elif result_type == "Hallucination": - hallucination += 1 - elif result_type == "Omission": - omission += 1 - - metrics = { - "correct_qa_ratio(all)": correct / total, - "hallucination_qa_ratio(all)": hallucination / total, - "omission_qa_ratio(all)": omission / total, - "qa_valid_num": valid, - "qa_num": total, - } - - if valid > 0: - metrics.update( - { - "correct_qa_ratio(valid)": correct / valid, - "hallucination_qa_ratio(valid)": hallucination / valid, - "omission_qa_ratio(valid)": omission / valid, - }, - ) - else: - metrics.update( - { - "correct_qa_ratio(valid)": 0, - "hallucination_qa_ratio(valid)": 0, - "omission_qa_ratio(valid)": 0, - }, - ) - - return metrics - - @staticmethod - def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: - """Compute question answering metrics for both result_type and original_result_type.""" - return { - "with_llm_answer": MetricsAggregator._compute_single_metric(qa_records, "result_type"), - "with_original_memories": MetricsAggregator._compute_single_metric(qa_records, "original_result_type"), - } - - @staticmethod - def compute_memory_integrity_metrics(integrity_records: list[dict]) -> dict[str, Any]: - """ - Compute memory integrity metrics. - - Args: - integrity_records: List of integrity evaluation results - - Returns: - dict with integrity metrics (score distribution and average) - """ - total = len(integrity_records) - if total == 0: - return { - "total_memory_points": 0, - "score_2_count": 0, - "score_1_count": 0, - "score_0_count": 0, - "score_2_ratio": 0, - "score_1_ratio": 0, - "score_0_ratio": 0, - "average_score": 0, - "valid_count": 0, - } - - score_2_count = 0 - score_1_count = 0 - score_0_count = 0 - valid_count = 0 - total_score = 0 - - for record in integrity_records: - score = record.get("integrity_score") - # Handle both string and int scores - if score is not None: - try: - score_int = int(score) - valid_count += 1 - total_score += score_int - if score_int == 2: - score_2_count += 1 - elif score_int == 1: - score_1_count += 1 - elif score_int == 0: - score_0_count += 1 - except (ValueError, TypeError): - pass - - metrics = { - "total_memory_points": total, - "score_2_count": score_2_count, - "score_1_count": score_1_count, - "score_0_count": score_0_count, - "score_2_ratio": score_2_count / total if total > 0 else 0, - "score_1_ratio": score_1_count / total if total > 0 else 0, - "score_0_ratio": score_0_count / total if total > 0 else 0, - "average_score": total_score / valid_count if valid_count > 0 else 0, - "accuracy": score_2_count / valid_count if valid_count > 0 else 0, - "valid_count": valid_count, - } - - return metrics - - @staticmethod - def compute_memory_accuracy_metrics(accuracy_records: list[dict]) -> dict[str, Any]: - """ - Compute memory accuracy metrics for extracted memories. - - Args: - accuracy_records: List of accuracy evaluation results - - Returns: - dict with accuracy metrics (score distribution, average, and inclusion ratio) - """ - total = len(accuracy_records) - if total == 0: - return { - "total_extracted_memories": 0, - "score_2_count": 0, - "score_1_count": 0, - "score_0_count": 0, - "score_2_ratio": 0, - "score_1_ratio": 0, - "score_0_ratio": 0, - "average_score": 0, - "accuracy": 0, - "included_in_golden_count": 0, - "included_in_golden_ratio": 0, - "valid_count": 0, - } - - score_2_count = 0 - score_1_count = 0 - score_0_count = 0 - included_count = 0 - valid_count = 0 - total_score = 0 - - for record in accuracy_records: - score = record.get("accuracy_score") - included = record.get("is_included_in_golden_memories") - - # Handle both string and int scores - if score is not None: - try: - score_int = int(score) - valid_count += 1 - total_score += score_int - if score_int == 2: - score_2_count += 1 - elif score_int == 1: - score_1_count += 1 - elif score_int == 0: - score_0_count += 1 - except (ValueError, TypeError): - pass - - # Handle is_included_in_golden_memories - if included is not None: - if isinstance(included, bool): - if included: - included_count += 1 - elif isinstance(included, str) and included.lower() == "true": - included_count += 1 - - metrics = { - "total_extracted_memories": total, - "score_2_count": score_2_count, - "score_1_count": score_1_count, - "score_0_count": score_0_count, - "score_2_ratio": score_2_count / total if total > 0 else 0, - "score_1_ratio": score_1_count / total if total > 0 else 0, - "score_0_ratio": score_0_count / total if total > 0 else 0, - "average_score": total_score / valid_count if valid_count > 0 else 0, - "accuracy": score_2_count / valid_count if valid_count > 0 else 0, - "included_in_golden_count": included_count, - "included_in_golden_ratio": included_count / total if total > 0 else 0, - "valid_count": valid_count, - } - - return metrics - - @staticmethod - def compute_time_metrics(eval_results_file: str) -> dict[str, float]: - """Compute timing metrics from evaluation results.""" - add_duration = 0 - search_duration = 0 - - with open(eval_results_file, "r", encoding="utf-8") as f: - for line in f: - if not line.strip(): - continue - user_data = json.loads(line) - - for session in user_data["sessions"]: - add_duration += session.get("add_dialogue_duration_ms", 0) - - eval_results = session.get("evaluation_results", {}) - for qa in eval_results.get("question_answering_records", []): - search_duration += qa.get("search_duration_ms", 0) - - # Convert to minutes - return { - "add_dialogue_duration_time": add_duration / 1000 / 60, - "search_memory_duration_time": search_duration / 1000 / 60, - "total_duration_time": (add_duration + search_duration) / 1000 / 60, - } - - -# ==================== Main Pipeline ==================== - - -class HaluMemEvaluator: - """HaluMem evaluator with proper resource management.""" - - def __init__(self, config: EvalConfig): - self.config = config - self.reme = ReMe( - default_llm_config={ - "model_name": self.config.reme_model_name, - }, - llms={ - "qwen-plus-t": { - "backend": "openai", - "model_name": "qwen-plus", - "extra_body": { - "enable_thinking": True, - }, - }, - "qwen-max-t": { - "backend": "openai", - "model_name": "qwen3-max", - "extra_body": { - "enable_thinking": True, - }, - }, - "gpt-4o-mini": { - "backend": "openai", - "model_name": "gpt-4o-mini-2024-07-18", - }, - "gpt-4o-mini-2024-07-18": { - "backend": "openai", - "model_name": "gpt-4o-mini-2024-07-18", - }, - "qwen-flash": { - "backend": "openai", - "model_name": "qwen-flash", - }, - }, - ) - - # Load evaluation prompts into ReMe's prompt handler - prompts_yaml_path = Path(__file__).parent / "eval_reme.yaml" - self.reme.prompt_handler.load_prompt_by_file(prompts_yaml_path) - - self.file_manager = FileManager(config.output_dir) - self.memory_processor = MemoryProcessor( - self.reme, - config.reme_model_name, - config.eval_model_name, - config.algo_version, - config.enable_thinking_params, - ) - self.qa_evaluator = QuestionAnsweringEvaluator( - self.memory_processor, - self.reme, - config.top_k, - config.eval_model_name, - ) - self.integrity_evaluator = MemoryIntegrityEvaluator( - self.reme, - eval_model_name="qwen-flash", - ) - self.accuracy_evaluator = MemoryAccuracyEvaluator( - self.reme, - eval_model_name="qwen-flash", - ) - self.data_loader = DataLoader() - - # For real-time updates - self._update_lock: asyncio.Lock | None = None - self._output_file: str | None = None - - async def __aenter__(self): - """Async context manager entry.""" - await self.reme.start() - return self - - async def __aexit__(self, exc_type, exc_val, exc_tb): - """Async context manager exit with cleanup.""" - await self.reme.close() - return False - - async def process_session( - self, - session: dict, - session_id: int, - user_name: str, - uuid: str, - ) -> dict: - """Process a single session using ReMe.""" - session_data = { - "uuid": uuid, - "user_name": user_name, - "session_id": session_id, - "memory_points": session["memory_points"], - } - - # Skip generated QA sessions - if session.get("is_generated_qa_session", False): - session_data["is_generated_qa_session"] = True - return session_data - - dialogue = session["dialogue"] - formatted_messages = self.data_loader.format_dialogue_messages(dialogue) - - extracted_memories, agent_messages, duration_ms = await self.memory_processor.add_memories( - user_id=user_name, - messages=formatted_messages, - batch_size=self.config.batch_size, - ) - - session_data.update( - { - "dialogue": dialogue, - "extracted_memories": extracted_memories, - "summary_messages": agent_messages, - "add_dialogue_duration_ms": duration_ms, - }, - ) - - # Evaluate memory integrity - check if extracted memories cover memory points - memory_points = session.get("memory_points", []) - formatted_dialogue = self.data_loader.format_dialogue_for_eval(dialogue, user_name) - - if memory_points and extracted_memories: - integrity_results = await self.integrity_evaluator.evaluate_memory_points( - extracted_memories=extracted_memories, - memory_points=memory_points, - ) - session_data["memory_integrity_results"] = integrity_results - - # Evaluate memory accuracy - check if each extracted memory is accurate - if extracted_memories and memory_points: - accuracy_results = await self.accuracy_evaluator.evaluate_extracted_memories( - extracted_memories=extracted_memories, - memory_points=memory_points, - formatted_dialogue=formatted_dialogue, - ) - session_data["memory_accuracy_results"] = accuracy_results - - # Evaluate questions if present - if "questions" in session: - formatted_dialogue = self.data_loader.format_dialogue_for_eval(dialogue, user_name) - qa_results = await self.qa_evaluator.evaluate_questions( - questions=session["questions"], - user_name=user_name, - uuid=uuid, - session_id=session_id, - formatted_dialogue=formatted_dialogue, - ) - - session_data["evaluation_results"] = { - "question_answering_records": qa_results, - } - - return session_data - - async def process_user(self, user_data: dict) -> dict: - """Process all sessions for a user.""" - user_name = self.data_loader.extract_user_name(user_data["persona_info"]) - uuid = user_data["uuid"] - - logger.info(f"Processing user: {user_name}") - - for idx, session in enumerate(user_data["sessions"]): - logger.info(f" Session {idx + 1}/{len(user_data['sessions'])}") - - session_data = await self.process_session( - session=session, - session_id=idx, - user_name=user_name, - uuid=uuid, - ) - - self.file_manager.save_session(user_name, idx, session_data) - - # Update results file after each session completes - await self._trigger_update() - - return {"uuid": uuid, "user_name": user_name, "status": "ok"} - - async def _trigger_update(self): - """Trigger real-time update of results and statistics.""" - if self._update_lock is None or self._output_file is None: - return - - async with self._update_lock: - self.file_manager.combine_results(self._output_file) - self._update_statistics(self._output_file) - - async def run_evaluation(self): - """Run the complete evaluation pipeline using ReMe.""" - start_time = time.time() - - # Load user data first to get user names - all_users = self.data_loader.load_jsonl(self.config.data_path) - users_to_process = all_users[: self.config.user_num] - - # Extract all user names and delete all profiles - all_user_names = [self.data_loader.extract_user_name(user_data["persona_info"]) for user_data in all_users] - if all_user_names: - for user_name in all_user_names: - self.reme.get_profile_handler(user_name).delete_all() - logger.info(f"Deleted all profiles for {len(all_user_names)} users") - - # Clear existing data - await self.reme.default_vector_store.delete_all() - - # Clear meta_memory directory - meta_memory_path = Path(f"meta_memory/{self.reme.default_vector_store.collection_name}") - if meta_memory_path.exists(): - shutil.rmtree(meta_memory_path) - logger.info(f"Cleared meta_memory directory: {meta_memory_path}") - meta_memory_path.mkdir(parents=True, exist_ok=True) - - print("\n" + "=" * 80) - print("HALUMEM EVALUATION - REME - QUESTION ANSWERING") - print(f"Users: {len(users_to_process)} | Concurrency: {self.config.max_concurrency}") - print("=" * 80 + "\n") - - # Output file path for real-time updates - self._output_file = os.path.join(self.config.output_dir, "eval_results.jsonl") - - # Lock for thread-safe file updates - self._update_lock = asyncio.Lock() - - # Process users with concurrency control - semaphore = asyncio.Semaphore(self.config.max_concurrency) - - async def process_with_cache_check(idx: int, user_data: dict): - async with semaphore: - user_name = self.data_loader.extract_user_name(user_data["persona_info"]) - - # Check cache - if self.file_manager.user_has_cache(user_name): - print(f"⚡ [{idx}/{len(users_to_process)}] Skipping {user_name} (cached)") - result = {"user_name": user_name, "status": "cached"} - # Also trigger update for cached users - await self._trigger_update() - else: - print(f"🔄 [{idx}/{len(users_to_process)}] Processing {user_name}...") - result = await self.process_user(user_data) - print(f"✅ [{idx}/{len(users_to_process)}] Completed {user_name}") - - return result - - tasks = [process_with_cache_check(idx, user) for idx, user in enumerate(users_to_process, 1)] - await asyncio.gather(*tasks) - - elapsed = time.time() - start_time - print(f"\n✅ Processing completed in {elapsed:.2f}s") - print(f"📁 Results: {self._output_file}\n") - - # Final aggregation and report - await self.aggregate_and_report(self._output_file) - - def _update_statistics(self, results_file: str): - """Update statistics file based on current results (for real-time monitoring).""" - if not os.path.exists(results_file): - return - - # Collect all QA records, memory integrity records, and accuracy records - qa_records = [] - integrity_records = [] - accuracy_records = [] - try: - with open(results_file, "r", encoding="utf-8") as f: - for line in f: - if not line.strip(): - continue - user_data = json.loads(line) - - for session in user_data["sessions"]: - if session.get("is_generated_qa_session"): - continue - - eval_results = session.get("evaluation_results", {}) - qa_records.extend( - eval_results.get("question_answering_records", []), - ) - - # Collect memory integrity records - integrity_records.extend( - session.get("memory_integrity_results", []), - ) - - # Collect memory accuracy records - accuracy_records.extend( - session.get("memory_accuracy_results", []), - ) - except (json.JSONDecodeError, KeyError): - return - - if not qa_records: - return - - # Compute metrics - qa_metrics = MetricsAggregator.compute_qa_metrics(qa_records) - time_metrics = MetricsAggregator.compute_time_metrics(results_file) - integrity_metrics = MetricsAggregator.compute_memory_integrity_metrics(integrity_records) - accuracy_metrics = MetricsAggregator.compute_memory_accuracy_metrics(accuracy_records) - - final_results = { - "overall_score": { - "question_answering": qa_metrics, - "memory_integrity": integrity_metrics, - "memory_accuracy": accuracy_metrics, - "time_consuming": time_metrics, - }, - "question_answering_records": qa_records, - "memory_integrity_records": integrity_records, - "memory_accuracy_records": accuracy_records, - } - - # Save statistics - report_file = os.path.join(self.config.output_dir, "eval_statistics.json") - with open(report_file, "w", encoding="utf-8") as f: - json.dump(final_results, f, ensure_ascii=False, indent=4) - - async def aggregate_and_report(self, results_file: str): - """Aggregate results and generate final report.""" - print("=" * 80) - print("AGGREGATING METRICS") - print("=" * 80 + "\n") - - # Collect all QA records - qa_records = [] - integrity_records = [] - accuracy_records = [] - with open(results_file, "r", encoding="utf-8") as f: - for line in f: - if not line.strip(): - continue - user_data = json.loads(line) - - for session in user_data["sessions"]: - if session.get("is_generated_qa_session"): - continue - - eval_results = session.get("evaluation_results", {}) - qa_records.extend( - eval_results.get("question_answering_records", []), - ) - - # Collect memory integrity records - integrity_records.extend( - session.get("memory_integrity_results", []), - ) - - # Collect memory accuracy records - accuracy_records.extend( - session.get("memory_accuracy_results", []), - ) - - # Compute metrics - qa_metrics = MetricsAggregator.compute_qa_metrics(qa_records) - time_metrics = MetricsAggregator.compute_time_metrics(results_file) - integrity_metrics = MetricsAggregator.compute_memory_integrity_metrics(integrity_records) - accuracy_metrics = MetricsAggregator.compute_memory_accuracy_metrics(accuracy_records) - - final_results = { - "overall_score": { - "question_answering": qa_metrics, - "memory_integrity": integrity_metrics, - "memory_accuracy": accuracy_metrics, - "time_consuming": time_metrics, - }, - "question_answering_records": qa_records, - "memory_integrity_records": integrity_records, - "memory_accuracy_records": accuracy_records, - } - - # Save final report - report_file = os.path.join(self.config.output_dir, "eval_statistics.json") - with open(report_file, "w", encoding="utf-8") as f: - json.dump(final_results, f, ensure_ascii=False, indent=4) - - print(f"📊 Statistics saved to: {report_file}\n") - - # Print summary - self._print_summary(qa_metrics, time_metrics, integrity_metrics, accuracy_metrics) - - def _print_summary( - self, - qa_metrics: dict, - time_metrics: dict, - integrity_metrics: dict = None, - accuracy_metrics: dict = None, - ): - """Print evaluation summary.""" - print("=" * 80) - print("EVALUATION SUMMARY - REME") - print("=" * 80 + "\n") - - # Print memory integrity metrics - if integrity_metrics and integrity_metrics.get("total_memory_points", 0) > 0: - print("🧠 Memory Integrity (coverage of expected memory points):") - total = integrity_metrics["total_memory_points"] - print( - f" Score 2 (Fully covered): " - f"{integrity_metrics['score_2_count']}/{total} " - f"({integrity_metrics['score_2_ratio']:.4f})", - ) - print( - f" Score 1 (Partially covered): " - f"{integrity_metrics['score_1_count']}/{total} " - f"({integrity_metrics['score_1_ratio']:.4f})", - ) - print( - f" Score 0 (Not covered): " - f"{integrity_metrics['score_0_count']}/{total} " - f"({integrity_metrics['score_0_ratio']:.4f})", - ) - print(f" Average Score: {integrity_metrics['average_score']:.4f}") - print(f" Accuracy (score=2 ratio): {integrity_metrics['accuracy']:.4f}") - print(f" Valid/Total: " f"{integrity_metrics['valid_count']}/{total}") - print() - - # Print memory accuracy metrics - if accuracy_metrics and accuracy_metrics.get("total_extracted_memories", 0) > 0: - print("🎯 Memory Accuracy (accuracy of extracted memories):") - total_acc = accuracy_metrics["total_extracted_memories"] - print( - f" Score 2 (Fully accurate): " - f"{accuracy_metrics['score_2_count']}/{total_acc} " - f"({accuracy_metrics['score_2_ratio']:.4f})", - ) - print( - f" Score 1 (Partially accurate): " - f"{accuracy_metrics['score_1_count']}/{total_acc} " - f"({accuracy_metrics['score_1_ratio']:.4f})", - ) - print( - f" Score 0 (Hallucinated): " - f"{accuracy_metrics['score_0_count']}/{total_acc} " - f"({accuracy_metrics['score_0_ratio']:.4f})", - ) - print(f" Average Score: {accuracy_metrics['average_score']:.4f}") - print(f" Accuracy (score=2 ratio): {accuracy_metrics['accuracy']:.4f}") - print( - f" Included in Golden: " - f"{accuracy_metrics['included_in_golden_count']}/{total_acc} " - f"({accuracy_metrics['included_in_golden_ratio']:.4f})", - ) - print(f" Valid/Total: " f"{accuracy_metrics['valid_count']}/{total_acc}") - print() - - # Print metrics for LLM-generated answer (result_type) - if qa_metrics and "with_llm_answer" in qa_metrics: - llm_metrics = qa_metrics["with_llm_answer"] - print("📊 Question Answering (with LLM answer):") - print(f" Correct (all): {llm_metrics['correct_qa_ratio(all)']:.4f}") - print(f" Hallucination (all): {llm_metrics['hallucination_qa_ratio(all)']:.4f}") - print(f" Omission (all): {llm_metrics['omission_qa_ratio(all)']:.4f}") - print(f" Correct (valid): {llm_metrics['correct_qa_ratio(valid)']:.4f}") - print(f" Hallucination (valid): {llm_metrics['hallucination_qa_ratio(valid)']:.4f}") - print(f" Omission (valid): {llm_metrics['omission_qa_ratio(valid)']:.4f}") - print(f" Valid/Total: {llm_metrics['qa_valid_num']}/{llm_metrics['qa_num']}") - - # Print metrics for original retrieved memories (original_result_type) - orig_metrics = qa_metrics["with_original_memories"] - print("\n📊 Question Answering (with original memories):") - print(f" Correct (all): {orig_metrics['correct_qa_ratio(all)']:.4f}") - print(f" Hallucination (all): {orig_metrics['hallucination_qa_ratio(all)']:.4f}") - print(f" Omission (all): {orig_metrics['omission_qa_ratio(all)']:.4f}") - print(f" Correct (valid): {orig_metrics['correct_qa_ratio(valid)']:.4f}") - print(f" Hallucination (valid): {orig_metrics['hallucination_qa_ratio(valid)']:.4f}") - print(f" Omission (valid): {orig_metrics['omission_qa_ratio(valid)']:.4f}") - print(f" Valid/Total: {orig_metrics['qa_valid_num']}/{orig_metrics['qa_num']}") - - print("\n⏱️ Time Metrics:") - print(f" Memory Addition: {time_metrics['add_dialogue_duration_time']:.2f} min") - print(f" Memory Search: {time_metrics['search_memory_duration_time']:.2f} min") - print(f" Total: {time_metrics['total_duration_time']:.2f} min") - print("\n" + "=" * 80) - - -# ==================== Entry Point ==================== - - -async def main_async( - data_path: str, - top_k: int, - batch_size: int, - user_num: int, - max_concurrency: int, - reme_model_name: str = "qwen-flash", - eval_model_name: str = "qwen3-max", - algo_version: str = "halumem", - enable_thinking_params: bool = False, -): - """Main async entry point for ReMe evaluation with proper resource cleanup.""" - config = EvalConfig( - data_path=data_path, - top_k=top_k, - batch_size=batch_size, - user_num=user_num, - max_concurrency=max_concurrency, - reme_model_name=reme_model_name, - eval_model_name=eval_model_name, - algo_version=algo_version, - enable_thinking_params=enable_thinking_params, - ) - - # Use async context manager for automatic cleanup - async with HaluMemEvaluator(config) as evaluator: - await evaluator.run_evaluation() - - -def main( - data_path: str, - top_k: int, - batch_size: int, - user_num: int, - max_concurrency: int, - reme_model_name: str = "qwen-flash", - eval_model_name: str = "qwen3-max", - algo_version: str = "halumem", - enable_thinking_params: bool = False, -): - """Main entry point for ReMe evaluation.""" - asyncio.run( - main_async( - data_path=data_path, - top_k=top_k, - batch_size=batch_size, - user_num=user_num, - max_concurrency=max_concurrency, - reme_model_name=reme_model_name, - eval_model_name=eval_model_name, - algo_version=algo_version, - enable_thinking_params=enable_thinking_params, - ), - ) - - -if __name__ == "__main__": - import argparse - - parser = argparse.ArgumentParser( - description="Evaluate ReMe on HaluMem benchmark (Question Answering)", - ) - parser.add_argument( - "--data_path", - type=str, - required=True, - help="Path to HaluMem JSONL file", - ) - parser.add_argument( - "--top_k", - type=int, - default=20, - help="Number of memories to retrieve (default: 20)", - ) - parser.add_argument( - "--user_num", - type=int, - default=1, - help="Number of users to evaluate (default: 1)", - ) - parser.add_argument( - "--max_concurrency", - type=int, - default=1, - help="Maximum concurrent user processing (default: 100)", - ) - parser.add_argument( - "--batch_size", - type=int, - default=40, - help="Batch size for memory summary processing of each conversation (default: 40)", - ) - parser.add_argument( - "--reme_model_name", - type=str, - default="qwen-flash", - help="Model name for ReMe (default: qwen-flash)", - ) - parser.add_argument( - "--eval_model_name", - type=str, - default="gpt-4o-mini-2024-07-18", - help="Model name for evaluation (default: qwen3-max)", - ) - parser.add_argument( - "--algo_version", - type=str, - default="default", - help="Algorithm version for summary and retrieval (default: default)", - ) - parser.add_argument( - "--enable_thinking_params", - action="store_true", - default=False, - help="Enable thinking parameters for summary and retrieval (default: False)", - ) - - args = parser.parse_args() - print(f"args={args}!") - - main( - data_path=args.data_path, - top_k=args.top_k, - batch_size=args.batch_size, - user_num=args.user_num, - max_concurrency=args.max_concurrency, - reme_model_name=args.reme_model_name, - eval_model_name=args.eval_model_name, - algo_version=args.algo_version, - enable_thinking_params=args.enable_thinking_params, - ) diff --git a/benchmark/locomo/eval_reme.py b/benchmark/locomo/eval_reme.py deleted file mode 100644 index bb424aee..00000000 --- a/benchmark/locomo/eval_reme.py +++ /dev/null @@ -1,1107 +0,0 @@ -""" -Simplified evaluation script for ReMe on Locomo benchmark. - -This script performs a simplified evaluation pipeline: -1. Load Locomo data -2. Process each user's sessions with ReMe (summary + retrieve) -3. Evaluate question answering -4. Generate metrics and statistics - -Usage: - python bench/halumem/eval_reme_simple.py --data_path locomo10.json \ - --top_k 20 --user_num 100 --max_concurrency 20 -""" - -import asyncio -import json -import os -import re -import shutil -import time -from pathlib import Path -from datetime import datetime, timezone, timedelta -from dataclasses import dataclass -from typing import Any -import yaml -from loguru import logger -from reme.core.enumeration import Role -from reme.core.schema import Message - - -from reme.reme import ReMe - - -# ==================== Configuration ==================== -@dataclass -class EvalConfig: - """Evaluation configuration parameters.""" - - data_path: str - top_k: int = 20 - user_num: int = 1 - max_concurrency: int = 2 - batch_size: int = 40 - output_dir: str = "bench_results/reme" - reme_model_name: str = "qwen-flash" - eval_model_name: str = "qwen3-max" - algo_version: str = "locomo" - enable_thinking_params: bool = False - - -# ==================== Utilities ==================== - - -class DataLoader: - """Handles loading and parsing of HaluMem data.""" - - @staticmethod - def load_jsonl(file_path: str) -> list[dict]: - """Load all entries from a JSONL file.""" - with open(file_path, "r", encoding="utf-8") as f: - return [json.loads(line.strip()) for line in f if line.strip()] - - @staticmethod - def load_json(file_path: str) -> dict: - """Load dict from a JSON file.""" - with open(file_path, "r", encoding="utf-8") as f: - return json.load(f) - - @staticmethod - def format_dialogue_messages( - dialogue: list[dict], - speaker_a: str, - base_timestamp: datetime, - time_interval: int, - ) -> list[dict]: - """Format dialogue into ReMe message format with conversation_time.""" - - return [ - { - "role": "user" if turn["speaker"] == speaker_a else "assistant", - "name": turn["speaker"], - "content": turn["text"], - "time_created": (base_timestamp + timedelta(seconds=idx * time_interval)).strftime("%Y-%m-%d %H:%M:%S"), - } - for idx, turn in enumerate(dialogue) - ] - - @staticmethod - def format_dialogue_for_eval(dialogue: list[dict], user_name: str = None) -> str: - """Format dialogue into string for evaluation.""" - formatted_turns = [] - for turn in dialogue: - timestamp = ( - datetime.strptime( - turn["timestamp"], - "%b %d, %Y, %H:%M:%S", - ) - .replace(tzinfo=timezone.utc) - .strftime("%Y-%m-%d %H:%M:%S") - ) - - # Use user_name if role is 'user' and user_name is provided - role = user_name if turn["role"] == "user" and user_name else turn["role"] - - formatted_turns.append( - f"Role: {role}\n" f"Content: {turn['content']}\n" f"Time: {timestamp}", - ) - return "\n\n".join(formatted_turns) - - -class FileManager: - """Manages file I/O operations.""" - - def __init__(self, base_dir: str): - self.base_dir = Path(base_dir) - self.tmp_dir = self.base_dir - self.tmp_dir.mkdir(parents=True, exist_ok=True) - - def get_user_dir(self, user_name: str) -> Path: - """Get the directory path for a user.""" - user_dir = self.tmp_dir / user_name - user_dir.mkdir(parents=True, exist_ok=True) - return user_dir - - def get_session_file(self, user_name: str, session_id: int) -> Path: - """Get the file path for a specific session.""" - return self.get_user_dir(user_name) / f"session_{session_id}.json" - - def get_question_file(self, user_name: str) -> Path: - """Get the file path for a specific question.""" - return self.get_user_dir(user_name) / "questions.json" - - def save_session(self, user_name: str, session_id: int, data: dict): - """Save session data to file.""" - file_path = self.get_session_file(user_name, session_id) - with open(file_path, "w", encoding="utf-8") as f: - json.dump(data, f, ensure_ascii=False, indent=2) - logger.info(f"✅ Saved session {session_id} to {file_path}") - - def save_question(self, user_name: str, data: dict): - """Save question data to file""" - file_path = self.get_question_file(user_name) - with open(file_path, "w", encoding="utf-8") as f: - json.dump(data, f, ensure_ascii=False, indent=2) - logger.info(f"✅ Saved question to {file_path}") - - def load_session(self, user_name: str, session_id: int) -> dict | None: - """Load session data from file.""" - file_path = self.get_session_file(user_name, session_id) - if not file_path.exists(): - return None - with open(file_path, "r", encoding="utf-8") as f: - return json.load(f) - - def user_has_cache(self, user_name: str) -> bool: - """Check if user has cached results.""" - user_dir = self.get_user_dir(user_name) - return any(f.name.startswith("session_") and f.suffix == ".json" for f in user_dir.iterdir()) - - def combine_results(self, output_file: str): - """Combine all user session files into a single JSONL file.""" - with open(output_file, "w", encoding="utf-8") as f_out: - for user_dir in self.tmp_dir.iterdir(): - if not user_dir.is_dir(): - continue - - session_files = sorted( - [f for f in user_dir.iterdir() if f.name.startswith("session_") and f.suffix == ".json"], - ) - - if not session_files: - continue - - # Load first session to get user metadata - with open(session_files[0], "r", encoding="utf-8") as f_in: - first_session = json.load(f_in) - - user_data = { - "uuid": first_session["uuid"], - "user_name": first_session["user_name"], - "sessions": [], - } - - # Load all sessions - for session_file in session_files: - with open(session_file, "r", encoding="utf-8") as f_in: - session_data = json.load(f_in) - # Remove redundant user metadata - session_data.pop("uuid", None) - session_data.pop("user_name", None) - user_data["sessions"].append(session_data) - - question_file = user_dir / "questions.json" - if not question_file.exists(): - continue - with open(question_file, "r", encoding="utf-8") as f_in: - question_data = json.load(f_in) - user_data["evaluation_results"] = { - "question_answering_records": question_data, - } - - f_out.write(json.dumps(user_data, ensure_ascii=False) + "\n") - - -# ==================== Memory Operations ==================== - - -class MemoryProcessor: - """Handles ReMe memory operations.""" - - def __init__( - self, - reme: ReMe, - eval_model_name: str = "qwen3-max", - algo_version: str = "locomo", - enable_thinking_params: bool = False, - ): - self.reme = reme - self.eval_model_name = eval_model_name - self.algo_version = algo_version - self.enable_thinking_params = enable_thinking_params - - async def add_memories( - self, - user_id: str, - messages: list[dict], - batch_size: int = 10000, - ) -> tuple[list[str], list, float]: - """ - Add memories in batches using ReMe and return extracted memory contents. - - Returns: - tuple: (extracted_memories, agent_messages, total_duration_ms) - """ - extracted_memories = [] - summary_messages = [] - total_duration_ms = 0 - - for i in range(0, len(messages), batch_size): - batch = messages[i : i + batch_size] - start = time.time() - - # Use new summary API - result = await self.reme.summarize_memory( - messages=batch, - user_name=user_id, - version=self.algo_version, - return_dict=True, - enable_time_filter=True, - enable_thinking_params=self.enable_thinking_params, - ) - - duration_ms = (time.time() - start) * 1000 - total_duration_ms += duration_ms - - extracted_memories.extend([m.model_dump(exclude_none=True) for m in result["answer"]]) - summary_messages.extend([m.simple_dump(enable_argument_dict=True) for m in result["messages"]]) - - return extracted_memories, summary_messages, total_duration_ms - - async def search_memory( - self, - query: str, - user_id: str, - top_k: int = 20, - ) -> tuple[dict, list, float]: - """ - Search memory using ReMe and return structured answer with reasoning. - - Returns: - tuple: (answer_dict, agent_messages, duration_ms) - answer_dict contains: {"reasoning": str, "answer": str, "memories": str} - """ - start = time.time() - - # Retrieve memories from ReMe using new API - result = await self.reme.retrieve_memory( - query=query, - retrieve_top_k=top_k, - user_name=user_id, - version=self.algo_version, - return_dict=True, - enable_time_filter=True, - enable_thinking_params=self.enable_thinking_params, - ) - - # Extract memories from response - memories = result["answer"] - agent_messages = [x.simple_dump(enable_argument_dict=True) for x in result["messages"]] - retrieved_nodes = [x.model_dump(exclude_none=True) for x in result["retrieved_nodes"]] - - # Use LLM to generate structured answer from memories - answer_result = await answer_question_with_memories( - reme=self.reme, - question=query, - memories=memories, - user_id=user_id, - model_name=self.eval_model_name, - ) - - # Add original memories to the result - answer_result["memories"] = memories - answer_result["retrieved_nodes"] = retrieved_nodes - - duration_ms = (time.time() - start) * 1000 - return answer_result, agent_messages, duration_ms - - -# ==================== Evaluation Functions ==================== - - -async def answer_question_with_memories( - reme: ReMe, - question: str, - memories: str, - user_id: str = None, - model_name: str = "qwen3-30b-a3b-instruct-2507", -): - """ - Answer a question using retrieved memories with PROMPT_MEMZERO_JSON template. - - Args: - reme: ReMe instance with default_llm and prompt_handler - question: The question to answer - memories: The retrieved memories (formatted as context) - user_id: Optional user ID for context formatting - model_name: Model name to use for LLM request - - Returns: - dict with 'reasoning' and 'answer' fields - """ - # Format context with memories - if user_id: - context = reme.prompt_handler.prompt_format( - "TEMPLATE_MEMOS", - user_id=user_id, - memories=memories, - ) - else: - context = f"Memories:\n{memories}" - - # Use PROMPT_MEMZERO_JSON template for structured JSON response - prompt = reme.prompt_handler.prompt_format( - "PROMPT_MEMZERO_JSON", - context=context, - question=question, - ) - - result = await reme.get_llm("qwen3_max_instruct").simple_request_for_json( - prompt=prompt, - model_name=model_name, - ) - - return result - - -async def evaluation_for_question( - reme: ReMe, - question: str, - golden_answer: str, - generated_answer: str, - model_name: str = "qwen3-max", -): - """ - Question-Answering Evaluation with optional Dialogue Context. - - Args: - reme: ReMe instance with default_llm and prompt_handler - question: The question string to be evaluated. - golden_answer: The reference (gold-standard) answer. - generated_answer: The answer produced by the memory system. - model_name: Model name to use for LLM request - - Returns: - dict with 'reasoning' and 'evaluation_result' fields - """ - await asyncio.sleep(10) - # Use configured prompts - system_prompt = reme.prompt_handler.prompt_format( - "SYSTEM_PROMPT", - ) - user_prompt = reme.prompt_handler.prompt_format( - "USER_PROMPT", - question=question, - golden_answer=golden_answer, - generated_answer=generated_answer, - ) - - reme_result = await reme.get_llm("qwen3_max_instruct").chat( - messages=[ - Message(role=Role.SYSTEM, content=system_prompt), - Message(role=Role.USER, content=user_prompt), - ], - model_name=model_name, - ) - - content = reme_result.content - match = re.search(r'"label"\s*:\s*"([^"]*?)"', content) - if match: - label = match.group(1) - else: - label = "WRONG" - result = { - "reasoning": content, - "evaluation_result": label.strip().upper() == "CORRECT", - } - return result - - -# ==================== Evaluation ==================== - - -class QuestionAnsweringEvaluator: - """Evaluates question answering performance.""" - - def __init__(self, memory_processor: MemoryProcessor, reme: ReMe, top_k: int, eval_model_name: str = "qwen3-max"): - self.memory_processor = memory_processor - self.reme = reme - self.top_k = top_k - self.eval_model_name = eval_model_name - - async def evaluate_questions( - self, - questions: list[dict], - user_name: str, - uuid: str, - ) -> list[dict]: - """Evaluate all questions for a conversation.""" - results = [] - - for qa in questions: - if qa["category"] == 5: - continue - answer_dict, agent_messages, duration_ms = await self.memory_processor.search_memory( - query=qa["question"], - user_id=user_name, - top_k=self.top_k, - ) - - # Extract answer and reasoning from the structured response - system_answer = answer_dict.get("answer", "") - system_reasoning = answer_dict.get("reasoning", "") - retrieved_memories = answer_dict.get("memories", "") - retrieved_nodes = answer_dict.get("retrieved_nodes", "") - - # Evaluate response - eval_result = await evaluation_for_question( - reme=self.reme, - question=qa["question"], - golden_answer=qa["answer"], - generated_answer=system_answer, - model_name=self.eval_model_name, - ) - - eval_result_original_answer = await evaluation_for_question( - reme=self.reme, - question=qa["question"], - golden_answer=qa["answer"], - generated_answer=retrieved_memories, - model_name=self.eval_model_name, - ) - - # Build result record - qa_result = { - **qa, - "uuid": uuid, - "system_response": system_answer, - "system_reasoning": system_reasoning, - "retrieved_memories": retrieved_memories, - "retrieved_nodes": retrieved_nodes, - "retrieve_messages": agent_messages, - "search_duration_ms": duration_ms, - "result_type": eval_result.get("evaluation_result"), - "question_answering_reasoning": eval_result.get("reasoning", ""), - "original_result_type": eval_result_original_answer.get("evaluation_result"), - "original_question_answering_reasoning": eval_result_original_answer.get("reasoning", ""), - } - results.append(qa_result) - - return results - - -class MetricsAggregator: - """Aggregates evaluation metrics.""" - - @staticmethod - def _compute_single_metric(qa_records: list[dict], result_key: str) -> dict[str, Any]: - """Compute metrics for a single result type key.""" - total = len(qa_records) - if total == 0: - return { - "correct_qa_ratio(all)": 0, - "correct_qa_ratio(valid)": 0, - "qa_valid_num": 0, - "qa_num": 0, - "category_1_accuracy": 0.0, - "category_2_accuracy": 0.0, - "category_3_accuracy": 0.0, - "category_4_accuracy": 0.0, - } - - correct = 0 - valid = 0 - - category_1_correct = 0 - category_1_num = 0 - category_1_valid = 0 - category_2_correct = 0 - category_2_num = 0 - category_2_valid = 0 - category_3_correct = 0 - category_3_num = 0 - category_3_valid = 0 - category_4_correct = 0 - category_4_num = 0 - category_4_valid = 0 - - for qa in qa_records: - result_type = qa.get(result_key, "") - category = qa.get("category", 0) - if category == 1: - category_1_num += 1 - elif category == 2: - category_2_num += 1 - elif category == 3: - category_3_num += 1 - elif category == 4: - category_4_num += 1 - - if result_type is not None and category in [1, 2, 3, 4]: - valid += 1 - if result_type is True: - correct += 1 - - if category == 1: - category_1_valid += 1 - if result_type is True: - category_1_correct += 1 - elif category == 2: - category_2_valid += 1 - if result_type is True: - category_2_correct += 1 - elif category == 3: - category_3_valid += 1 - if result_type is True: - category_3_correct += 1 - elif category == 4: - category_4_valid += 1 - if result_type is True: - category_4_correct += 1 - - metrics = { - "correct_qa_ratio(all)": correct / total, - "qa_valid_num": valid, - "qa_num": total, - "category_1_accuracy": category_1_correct / category_1_num if category_1_num > 0 else 0, - "category_1_num": category_1_num, - "category_1_valid_num": category_1_valid, - "category_2_accuracy": category_2_correct / category_2_num if category_2_num > 0 else 0, - "category_2_num": category_2_num, - "category_2_valid_num": category_2_valid, - "category_3_accuracy": category_3_correct / category_3_num if category_3_num > 0 else 0, - "category_3_num": category_3_num, - "category_3_valid_num": category_3_valid, - "category_4_accuracy": category_4_correct / category_4_num if category_4_num > 0 else 0, - "category_4_num": category_4_num, - "category_4_valid_num": category_4_valid, - } - - if valid > 0: - metrics.update( - { - "correct_qa_ratio(valid)": correct / valid, - }, - ) - else: - metrics.update( - { - "correct_qa_ratio(valid)": 0, - }, - ) - - return metrics - - @staticmethod - def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]: - """Compute question answering metrics for both result_type and original_result_type.""" - return { - "with_llm_answer": MetricsAggregator._compute_single_metric(qa_records, "result_type"), - "with_original_memories": MetricsAggregator._compute_single_metric(qa_records, "original_result_type"), - } - - @staticmethod - def compute_time_metrics(eval_results_file: str) -> dict[str, float]: - """Compute timing metrics from evaluation results.""" - add_duration = 0 - search_duration = 0 - - with open(eval_results_file, "r", encoding="utf-8") as f: - for line in f: - if not line.strip(): - continue - user_data = json.loads(line) - - for session in user_data["sessions"]: - add_duration += session.get("add_dialogue_duration_ms", 0) - - eval_results = user_data.get("evaluation_results", {}) - for qa in eval_results.get("question_answering_records", []): - search_duration += qa.get("search_duration_ms", 0) - - # Convert to minutes - return { - "add_dialogue_duration_time": add_duration / 1000 / 60, - "search_memory_duration_time": search_duration / 1000 / 60, - "total_duration_time": (add_duration + search_duration) / 1000 / 60, - } - - -# ==================== Evaluator ==================== - - -class LocomoEvaluator: - """ - LOCOMO 评估器核心类 - 用于评估 MemAgent 的记忆完整性、记忆准确性和问答准确性 - """ - - def __init__(self, config: EvalConfig): - self.config = config - with open("eval_reme.yaml", "r", encoding="utf-8") as file: - data = yaml.safe_load(file) - self.summary_prompt_1 = data["user_message_summary_1"] - self.summary_prompt_2 = data["user_message_summary_2"] - self.retriever_prompt = data["user_message_retrieve"] - - ops_dict = { - "personal_summarizer": { - "prompt_dict": { - "user_message_s1": self.summary_prompt_1, - "user_message_s2": self.summary_prompt_2, - }, - }, - "personal_retriever": { - "prompt_dict": { - "user_message_s2": self.retriever_prompt, - }, - "params": { - "return_memory_nodes": True, - }, - }, - } - - self.reme = ReMe( - default_llm_config={ - "model_name": self.config.reme_model_name, - }, - ops=ops_dict, - ) - - # Load evaluation prompts into ReMe's prompt handler - prompts_yaml_path = Path(__file__).parent / "eval_reme.yaml" - self.reme.prompt_handler.load_prompt_by_file(prompts_yaml_path) - - self.file_manager = FileManager(config.output_dir) - self.memory_processor = MemoryProcessor( - self.reme, - config.eval_model_name, - config.algo_version, - config.enable_thinking_params, - ) - self.qa_evaluator = QuestionAnsweringEvaluator( - self.memory_processor, - self.reme, - config.top_k, - config.eval_model_name, - ) - self.data_loader = DataLoader() - - # For real-time updates - self._update_lock: asyncio.Lock | None = None - self._output_file: str | None = None - - async def __aenter__(self): - """Async context manager entry.""" - await self.reme.start() - return self - - async def __aexit__(self, exc_type, exc_val, exc_tb): - """Async context manager exit with cleanup.""" - await self.reme.close() - return False - - async def process_user(self, user_data: dict) -> dict: - """Process all sessions for a user.""" - speaker_a = user_data["conversation"]["speaker_a"] - speaker_b = user_data["conversation"]["speaker_b"] - uuid = f"{speaker_a}_{speaker_b}" - user_name = [speaker_a, speaker_b] - user_file_name = f"{speaker_a}_{speaker_b}" - - new_user_data = { - "uuid": f"{user_data['conversation']['speaker_a']}_{user_data['conversation']['speaker_b']}", - "user_name": user_name, - "sessions": [], - "qas": [], - "eval_results": {}, - } - logger.info(f"Processing user: {speaker_a} and {speaker_b}") - session_num = 19 if uuid == "Caroline_Melanie" else int(len(user_data["conversation"]) / 2 - 1) - time_interval = 60 - - # Process conversation - for idx in range(session_num): - conversation = user_data["conversation"] - logger.info(f"Processing user {user_name}: session {idx+1}/{session_num}") - session_data = { - "uuid": uuid, - "user_name": user_file_name, - "timestamp": conversation[f"session_{idx+1}_date_time"], - "session": conversation[f"session_{idx+1}"], - } - - # Format dialogue - dialogue = conversation[f"session_{idx+1}"] - base_timestamp = parse_locomo_timestamp(session_data["timestamp"]) - formatted_messages = self.data_loader.format_dialogue_messages( - dialogue, - speaker_a, - base_timestamp, - time_interval, - ) - extracted_memories, agent_messages, duration_ms = await self.memory_processor.add_memories( - user_id=user_name, - messages=formatted_messages, - batch_size=self.config.batch_size, - ) - session_data.update( - { - "dialogue": dialogue, - "extracted_memories": extracted_memories, - "summary_messages": agent_messages, - "add_dialogue_duration_ms": duration_ms, - }, - ) - - self.file_manager.save_session(user_file_name, idx, session_data) - - # Process questions - qas = user_data["qa"] - qa_results = await self.qa_evaluator.evaluate_questions( - questions=qas, - user_name=user_name, - uuid=uuid, - ) - - new_user_data["evaluation_results"] = { - "question_answering_records": qa_results, - } - self.file_manager.save_question(user_file_name, qa_results) - - # Update results file after each conversation completes - await self._trigger_update() - - return {"uuid": uuid, "user_name": user_name, "status": "ok"} - - async def _trigger_update(self): - """Trigger real-time update of results and statistics.""" - if self._update_lock is None or self._output_file is None: - return - - async with self._update_lock: - self.file_manager.combine_results(self._output_file) - self._update_statistics(self._output_file) - - async def run_evaluation(self): - """Run the complete evaluation pipeline using ReMe.""" - start_time = time.time() - - # Load user data first to get user names - all_users = self.data_loader.load_json(self.config.data_path) - users_to_process = all_users[: self.config.user_num] - - # Extract all user names and delete all profiles - all_user_names = [ - f"{user_data['conversation']['speaker_a']}_&_{user_data['conversation']['speaker_b']}" - for user_data in all_users - ] - if all_user_names: - for user_name in all_user_names: - self.reme.get_profile_handler(user_name).delete_all() - logger.info(f"Deleted all profiles for {len(all_user_names)} users") - - # Clear existing data - await self.reme.default_vector_store.delete_all() - - # Clear meta_memory directory - meta_memory_path = Path(f"meta_memory/{self.reme.default_vector_store.collection_name}") - if meta_memory_path.exists(): - shutil.rmtree(meta_memory_path) - logger.info(f"Cleared meta_memory directory: {meta_memory_path}") - meta_memory_path.mkdir(parents=True, exist_ok=True) - - print("\n" + "=" * 80) - print("LOCOMO EVALUATION - REME - QUESTION ANSWERING") - print(f"Users: {len(users_to_process)} | Concurrency: {self.config.max_concurrency}") - print("=" * 80 + "\n") - - # Output file path for real-time updates - self._output_file = os.path.join(self.config.output_dir, "eval_results.jsonl") - - # Lock for thread-safe file updates - self._update_lock = asyncio.Lock() - - # Process users with concurrency control - semaphore = asyncio.Semaphore(self.config.max_concurrency) - - async def process_with_cache_check(idx: int, user_data: dict): - async with semaphore: - user_name = f"{user_data['conversation']['speaker_a']}_{user_data['conversation']['speaker_b']}" - - # Check cache - if self.file_manager.user_has_cache(user_name): - print(f"⚡ [{idx}/{len(users_to_process)}] Skipping {user_name} (cached)") - result = {"user_name": user_name, "status": "cached"} - # Also trigger update for cached users - await self._trigger_update() - else: - print(f"🔄 [{idx}/{len(users_to_process)}] Processing {user_name}...") - result = await self.process_user(user_data) - print(f"✅ [{idx}/{len(users_to_process)}] Completed {user_name}") - - return result - - tasks = [process_with_cache_check(idx, user) for idx, user in enumerate(users_to_process, 1)] - await asyncio.gather(*tasks, return_exceptions=True) - - elapsed = time.time() - start_time - print(f"\n✅ Processing completed in {elapsed:.2f}s") - print(f"📁 Results: {self._output_file}\n") - - # Final aggregation and report - await self.aggregate_and_report(self._output_file) - - def _update_statistics(self, results_file: str): - """Update statistics file based on current results (for real-time monitoring).""" - if not os.path.exists(results_file): - return - - # Collect all QA records - qa_records = [] - try: - with open(results_file, "r", encoding="utf-8") as f: - for line in f: - if not line.strip(): - continue - user_data = json.loads(line) - eval_results = user_data.get("evaluation_results", {}) - qa_records.extend( - eval_results.get("question_answering_records", []), - ) - except (json.JSONDecodeError, KeyError): - return - - if not qa_records: - return - - # Compute metrics - qa_metrics = MetricsAggregator.compute_qa_metrics(qa_records) - time_metrics = MetricsAggregator.compute_time_metrics(results_file) - - final_results = { - "overall_score": { - "question_answering": qa_metrics, - "time_consuming": time_metrics, - }, - "question_answering_records": qa_records, - } - - # Save statistics - report_file = os.path.join(self.config.output_dir, "eval_statistics.json") - with open(report_file, "w", encoding="utf-8") as f: - json.dump(final_results, f, ensure_ascii=False, indent=4) - - async def aggregate_and_report(self, results_file: str): - """Aggregate results and generate final report.""" - print("=" * 80) - print("AGGREGATING METRICS") - print("=" * 80 + "\n") - - # Collect all QA records - qa_records = [] - print(results_file) - with open(results_file, "r", encoding="utf-8") as f: - for line in f: - if not line.strip(): - continue - user_data = json.loads(line) - print(user_data) - eval_results = user_data.get("evaluation_results", {}) - qa_records.extend( - eval_results.get("question_answering_records", []), - ) - - # Compute metrics - qa_metrics = MetricsAggregator.compute_qa_metrics(qa_records) - time_metrics = MetricsAggregator.compute_time_metrics(results_file) - - final_results = { - "overall_score": { - "question_answering": qa_metrics, - "time_consuming": time_metrics, - }, - "question_answering_records": qa_records, - } - - # Save final report - report_file = os.path.join(self.config.output_dir, "eval_statistics.json") - with open(report_file, "w", encoding="utf-8") as f: - json.dump(final_results, f, ensure_ascii=False, indent=4) - - print(f"📊 Statistics saved to: {report_file}\n") - - # Print summary - self._print_summary(qa_metrics, time_metrics) - - def _print_summary(self, qa_metrics: dict, time_metrics: dict): - """Print evaluation summary.""" - print("=" * 80) - print("EVALUATION SUMMARY - REME") - print("=" * 80 + "\n") - - # Print metrics for LLM-generated answer (result_type) - llm_metrics = qa_metrics["with_llm_answer"] - print(llm_metrics) - print("📊 Question Answering (with LLM answer):") - print(f" Correct (all): {llm_metrics['correct_qa_ratio(all)']:.4f}") - print(f" Correct (valid): {llm_metrics['correct_qa_ratio(valid)']:.4f}") - print(f" Valid/Total: {llm_metrics['qa_valid_num']}/{llm_metrics['qa_num']}") - print(f" Category 1 Accuracy: {llm_metrics['category_1_accuracy']:.4f}") - print(f" Category 2 Accuracy: {llm_metrics['category_2_accuracy']:.4f}") - print(f" Category 3 Accuracy: {llm_metrics['category_3_accuracy']:.4f}") - print(f" Category 4 Accuracy: {llm_metrics['category_4_accuracy']:.4f}") - - # Print metrics for original retrieved memories (original_result_type) - orig_metrics = qa_metrics["with_original_memories"] - print("\n📊 Question Answering (with original memories):") - print(f" Correct (all): {orig_metrics['correct_qa_ratio(all)']:.4f}") - print(f" Correct (valid): {orig_metrics['correct_qa_ratio(valid)']:.4f}") - print(f" Valid/Total: {orig_metrics['qa_valid_num']}/{orig_metrics['qa_num']}") - print(f" Category 1 Accuracy: {orig_metrics['category_1_accuracy']:.4f}") - print(f" Category 2 Accuracy: {orig_metrics['category_2_accuracy']:.4f}") - print(f" Category 3 Accuracy: {orig_metrics['category_3_accuracy']:.4f}") - print(f" Category 4 Accuracy: {orig_metrics['category_4_accuracy']:.4f}") - - print("\n⏱️ Time Metrics:") - print(f" Memory Addition: {time_metrics['add_dialogue_duration_time']:.2f} min") - print(f" Memory Search: {time_metrics['search_memory_duration_time']:.2f} min") - print(f" Total: {time_metrics['total_duration_time']:.2f} min") - print("\n" + "=" * 80) - - -def parse_locomo_timestamp(timestamp_str: str): - """ - Parse LoCoMo timestamp format. - - Input format: "6:07 pm on 13 January, 2023" - Special value: "Unknown" or unparseable returns None - Output: datetime object or None - """ - # Clean string - timestamp_str = timestamp_str.replace("\\s+", " ").strip() - - # Handle special cases: Unknown or empty string - if timestamp_str.lower() == "unknown" or not timestamp_str: - # No time information, return None - return None - - try: - return datetime.strptime(timestamp_str, "%I:%M %p on %d %B, %Y") - except ValueError: - # If parse fails, return None and print warning - print(f"⚠️ Warning: Failed to parse timestamp '{timestamp_str}', no timestamp will be set") - return None - - -# ==================== Main Pipeline ==================== - - -async def main_async( - data_path: str, - top_k: int, - user_num: int, - max_concurrency: int, - reme_model_name: str = "qwen-flash", - eval_model_name: str = "qwen3-max", - algo_version: str = "halumem", - enable_thinking_params: bool = False, -): - """Main async entry point for ReMe evaluation with proper resource cleanup.""" - config = EvalConfig( - data_path=data_path, - top_k=top_k, - user_num=user_num, - max_concurrency=max_concurrency, - reme_model_name=reme_model_name, - eval_model_name=eval_model_name, - algo_version=algo_version, - enable_thinking_params=enable_thinking_params, - ) - - # Use async context manager for automatic cleanup - async with LocomoEvaluator(config) as evaluator: - await evaluator.run_evaluation() - - -def main( - data_path: str, - top_k: int = 20, - user_num: int = 1, - max_concurrency: int = 2, - reme_model_name: str = "qwen-flash", - eval_model_name: str = "qwen3-max", - algo_version: str = "halumem", - enable_thinking_params: bool = False, -): - """Synchronous entry point.""" - asyncio.run( - main_async( - data_path=data_path, - top_k=top_k, - user_num=user_num, - max_concurrency=max_concurrency, - reme_model_name=reme_model_name, - eval_model_name=eval_model_name, - algo_version=algo_version, - enable_thinking_params=enable_thinking_params, - ), - ) - - -if __name__ == "__main__": - import argparse - - parser = argparse.ArgumentParser(description="Simplified evaluation for ReMe on Locomo benchmark") - parser.add_argument( - "--data_path", - type=str, - required=True, - help="Path to Locomo data file (e.g., locomo10.jsonl)", - ) - parser.add_argument( - "--top_k", - type=int, - default=20, - help="Number of top memories to retrieve (default: 20)", - ) - parser.add_argument( - "--user_num", - type=int, - default=1, - help="Number of users to evaluate (default: 1)", - ) - parser.add_argument( - "--max_concurrency", - type=int, - default=2, - help="Maximum concurrency for processing (default: 2)", - ) - parser.add_argument( - "--reme_model_name", - type=str, - default="qwen-flash", - help="Model name for ReMe (default: qwen-flash)", - ) - parser.add_argument( - "--eval_model_name", - type=str, - default="qwen3-max", - help="Model name for evaluation (default: qwen3-max)", - ) - parser.add_argument( - "--algo_version", - type=str, - default="default", - help="Algorithm version for summary and retrieval (default: halumem)", - ) - parser.add_argument( - "--enable_thinking_params", - action="store_true", - default=True, - help="Enable thinking parameters for summary and retrieval (default: False)", - ) - - args = parser.parse_args() - print(f"args={args}!") - - main( - data_path=args.data_path, - top_k=args.top_k, - user_num=args.user_num, - max_concurrency=args.max_concurrency, - reme_model_name=args.reme_model_name, - eval_model_name=args.eval_model_name, - algo_version=args.algo_version, - enable_thinking_params=args.enable_thinking_params, - ) diff --git a/benchmark/longmemeval/compute_stats.py b/benchmark/longmemeval/compute_stats.py deleted file mode 100644 index 3e500f1f..00000000 --- a/benchmark/longmemeval/compute_stats.py +++ /dev/null @@ -1,346 +0,0 @@ -""" -LongMemEval Evaluation Statistics Analyzer - -Computes detailed statistics from evaluation results including: -- Overall accuracy -- Accuracy by question type -- Timing statistics (summary, retrieval) -- Memory extraction statistics - -Usage: - python bench/longmemeval/compute_stats.py \ - --results_dir bench/longmemeval/bench_results/longmemeval_reme -""" - -import argparse -import json -from collections import defaultdict -from pathlib import Path -from typing import Any - - -def load_results(results_dir: str) -> list[dict]: - """Load all question result files from the directory. - - Args: - results_dir: Path to the results directory - - Returns: - List of result dictionaries - """ - results_path = Path(results_dir) - results = [] - - # Load individual question files - question_files = sorted(results_path.glob("question_*.json")) - - for file_path in question_files: - try: - with open(file_path, "r", encoding="utf-8") as f: - result = json.load(f) - results.append(result) - except Exception as e: - print(f"⚠️ Error loading {file_path}: {e}") - - return results - - -def compute_accuracy_stats(results: list[dict]) -> dict[str, Any]: - """Compute overall and per-type accuracy statistics. - - Args: - results: List of result dictionaries - - Returns: - Dictionary with accuracy statistics - """ - total = len(results) - correct = 0 - incorrect = 0 - error = 0 - - # Per question type statistics - type_stats = defaultdict(lambda: {"total": 0, "correct": 0, "incorrect": 0, "error": 0}) - - for r in results: - qtype = r.get("question_type", "unknown") - judgment = r.get("judgment", {}) - is_correct = judgment.get("is_correct") - - type_stats[qtype]["total"] += 1 - - if is_correct is True: - correct += 1 - type_stats[qtype]["correct"] += 1 - elif is_correct is False: - incorrect += 1 - type_stats[qtype]["incorrect"] += 1 - else: - error += 1 - type_stats[qtype]["error"] += 1 - - # Compute accuracies - overall = { - "total": total, - "correct": correct, - "incorrect": incorrect, - "error": error, - "accuracy": correct / total if total > 0 else 0, - "accuracy_valid": correct / (correct + incorrect) if (correct + incorrect) > 0 else 0, - } - - by_type = {} - for qtype, stats in type_stats.items(): - valid = stats["correct"] + stats["incorrect"] - by_type[qtype] = { - **stats, - "accuracy": stats["correct"] / stats["total"] if stats["total"] > 0 else 0, - "accuracy_valid": stats["correct"] / valid if valid > 0 else 0, - } - - return { - "overall": overall, - "by_question_type": by_type, - } - - -def compute_timing_stats(results: list[dict]) -> dict[str, Any]: - """Compute timing statistics. - - Args: - results: List of result dictionaries - - Returns: - Dictionary with timing statistics - """ - summary_times = [] - retrieve_times = [] - - for r in results: - summary_ms = r.get("summary_duration_ms", 0) - retrieve_ms = r.get("retrieve_duration_ms", 0) - - if summary_ms > 0: - summary_times.append(summary_ms) - if retrieve_ms > 0: - retrieve_times.append(retrieve_ms) - - def compute_stats(times: list[float]) -> dict: - if not times: - return {"count": 0, "total_ms": 0, "avg_ms": 0, "min_ms": 0, "max_ms": 0} - - return { - "count": len(times), - "total_ms": sum(times), - "total_min": sum(times) / 1000 / 60, - "avg_ms": sum(times) / len(times), - "min_ms": min(times), - "max_ms": max(times), - } - - return { - "summary": compute_stats(summary_times), - "retrieve": compute_stats(retrieve_times), - "total_time_min": (sum(summary_times) + sum(retrieve_times)) / 1000 / 60, - } - - -def compute_memory_stats(results: list[dict]) -> dict[str, Any]: - """Compute memory extraction statistics. - - Args: - results: List of result dictionaries - - Returns: - Dictionary with memory statistics - """ - memory_counts = [] - session_counts = [] - - for r in results: - memories = r.get("extracted_memories", []) - num_sessions = r.get("num_sessions", 0) - - memory_counts.append(len(memories)) - session_counts.append(num_sessions) - - def compute_stats(counts: list[int]) -> dict: - if not counts: - return {"count": 0, "total": 0, "avg": 0, "min": 0, "max": 0} - - return { - "count": len(counts), - "total": sum(counts), - "avg": sum(counts) / len(counts), - "min": min(counts), - "max": max(counts), - } - - return { - "memories_per_question": compute_stats(memory_counts), - "sessions_per_question": compute_stats(session_counts), - } - - -def print_report( - accuracy_stats: dict, - timing_stats: dict, - memory_stats: dict, - results_dir: str, -): - """Print formatted statistics report. - - Args: - accuracy_stats: Accuracy statistics - timing_stats: Timing statistics - memory_stats: Memory statistics - results_dir: Path to results directory - """ - print("\n" + "=" * 80) - print("LONGMEMEVAL EVALUATION STATISTICS") - print(f"Results Directory: {results_dir}") - print("=" * 80) - - # Overall accuracy - overall = accuracy_stats["overall"] - print("\n📊 Overall Accuracy:") - print(f" Total Questions: {overall['total']}") - print(f" ✅ Correct: {overall['correct']} ({100 * overall['accuracy']:.2f}%)") - print( - f" ❌ Incorrect: {overall['incorrect']} " - f"({100 * overall['incorrect'] / overall['total'] if overall['total'] > 0 else 0:.2f}%)", - ) - if overall["error"] > 0: - print(f" ⚠️ Error: {overall['error']} ({100 * overall['error'] / overall['total']:.2f}%)") - print(f" Accuracy (valid): {100 * overall['accuracy_valid']:.2f}%") - - # Accuracy by question type - print("\n📊 Accuracy by Question Type:") - print("-" * 60) - print(f"{'Question Type':<30} {'Correct':<10} {'Total':<10} {'Accuracy':<10}") - print("-" * 60) - - by_type = accuracy_stats["by_question_type"] - for qtype in sorted(by_type.keys()): - stats = by_type[qtype] - print(f"{qtype:<30} {stats['correct']:<10} {stats['total']:<10} {100 * stats['accuracy']:.2f}%") - - print("-" * 60) - - # Timing statistics - print("\n⏱️ Timing Statistics:") - summary = timing_stats["summary"] - retrieve = timing_stats["retrieve"] - - print(" Memory Summarization:") - print(f" Total Time: {summary['total_min']:.2f} min") - print(f" Avg per Q: {summary['avg_ms']:.0f} ms") - print(f" Min/Max: {summary['min_ms']:.0f} / {summary['max_ms']:.0f} ms") - - print(" Memory Retrieval:") - print(f" Total Time: {retrieve['total_min']:.2f} min") - print(f" Avg per Q: {retrieve['avg_ms']:.0f} ms") - print(f" Min/Max: {retrieve['min_ms']:.0f} / {retrieve['max_ms']:.0f} ms") - - print(f" Total Time: {timing_stats['total_time_min']:.2f} min") - - # Memory statistics - print("\n📝 Memory Statistics:") - mem = memory_stats["memories_per_question"] - sess = memory_stats["sessions_per_question"] - - print(" Extracted Memories per Question:") - print(f" Total: {mem['total']}") - print(f" Average: {mem['avg']:.1f}") - print(f" Min/Max: {mem['min']} / {mem['max']}") - - print(" Sessions per Question:") - print(f" Average: {sess['avg']:.1f}") - print(f" Min/Max: {sess['min']} / {sess['max']}") - - print("\n" + "=" * 80) - - -def save_statistics( - accuracy_stats: dict, - timing_stats: dict, - memory_stats: dict, - output_file: str, -): - """Save statistics to JSON file. - - Args: - accuracy_stats: Accuracy statistics - timing_stats: Timing statistics - memory_stats: Memory statistics - output_file: Path to output file - """ - stats = { - "accuracy": accuracy_stats, - "timing": timing_stats, - "memory": memory_stats, - } - - with open(output_file, "w", encoding="utf-8") as f: - json.dump(stats, f, indent=4, ensure_ascii=False) - - print(f"\n📁 Statistics saved to: {output_file}") - - -def main(results_dir: str, output_file: str = None): - """Main function to compute and display statistics. - - Args: - results_dir: Path to results directory - output_file: Optional path to save statistics JSON - """ - print(f"\nLoading results from: {results_dir}") - - results = load_results(results_dir) - - if not results: - print("❌ No results found!") - return - - print(f"Loaded {len(results)} question results") - - # Compute statistics - accuracy_stats = compute_accuracy_stats(results) - timing_stats = compute_timing_stats(results) - memory_stats = compute_memory_stats(results) - - # Print report - print_report(accuracy_stats, timing_stats, memory_stats, results_dir) - - # Save to file if specified - if output_file: - save_statistics(accuracy_stats, timing_stats, memory_stats, output_file) - else: - # Default output file in results directory - default_output = Path(results_dir) / "statistics.json" - save_statistics(accuracy_stats, timing_stats, memory_stats, str(default_output)) - - -if __name__ == "__main__": - parser = argparse.ArgumentParser( - description="Compute statistics from LongMemEval evaluation results", - ) - parser.add_argument( - "--results_dir", - type=str, - default="bench_results/longmemeval_reme", - help="Path to results directory containing question_*.json files", - ) - parser.add_argument( - "--output_file", - type=str, - default=None, - help="Path to save statistics JSON (default: /statistics.json)", - ) - - args = parser.parse_args() - - main( - results_dir=args.results_dir, - output_file=args.output_file, - ) diff --git a/benchmark/longmemeval/eval_longmemeval_reme.py b/benchmark/longmemeval/eval_longmemeval_reme.py deleted file mode 100644 index c529b434..00000000 --- a/benchmark/longmemeval/eval_longmemeval_reme.py +++ /dev/null @@ -1,1087 +0,0 @@ -""" -LongMemEval Benchmark Evaluator for ReMe - -A modular evaluation pipeline that: -1. Loads LongMemEval benchmark data (each entry is a question with haystack sessions) -2. Processes haystack sessions through ReMe for memory summarization -3. Uses questions to query memory and generate answers -4. Uses LLM to judge answer correctness -5. Generates comprehensive metrics - -Usage: - python benchmark/longmemeval/eval_longmemeval_reme.py \ - --data_path dataset/longmemeval/longmemeval_s_cleaned.json \ - --top_k 20 --start_index 0 --end_index 10 -""" - -import asyncio -import json -import time -from dataclasses import dataclass -from datetime import datetime, timezone, timedelta -from pathlib import Path -from typing import Any, Optional - -from loguru import logger - -from reme.reme import ReMe - - -# ==================== Configuration ==================== - - -@dataclass -class EvalConfig: - """Evaluation configuration parameters.""" - - data_path: str - top_k: int = 10 - start_index: int = 0 - end_index: Optional[int] = None - max_concurrency: int = 1 - batch_size: int = 30 - output_dir: str = "cache/bench_results/longmemeval_reme" - reme_model_name: str = "qwen-flash" # summary模型 - retrieve_model_name: str = "qwen-max" # retrieve模型 - eval_model_name: str = "qwen-max" # 评估/判断模型 - algo_version: str = "v1" - samples_per_type: int = -1 # Number of samples per question type, -1 for all - enable_thinking_params: bool = False - - -# ==================== Answer Judge Prompts ==================== - - -def get_anscheck_prompt(task: str, question: str, answer: str, response: str, abstention: bool = False) -> str: - """Generate the answer checking prompt based on question type. - - Args: - task: Question type, e.g. 'single-session-user', 'multi-session', 'temporal-reasoning' - question: The question content - answer: The reference answer - response: The model's response - abstention: Whether this is an unanswerable question - - Returns: - Prompt for judging answer correctness - """ - if not abstention: - if task in ["single-session-user", "single-session-assistant", "multi-session"]: - template = ( - "I will give you a question, a correct answer, and a response from a model. Please answer yes " - "if the response contains the correct answer. Otherwise, answer no. If the response is equi" - "valent to the correct answer or contains all the intermediate steps to get the correct answer," - " you should also answer yes. If the response only contains a subset of the information requir" - "ed by the answer, answer no. \n\nQuestion: {}\n\nCorrect Answer: {}\n\nModel Response: " - "{}\n\nIs the model response correct? Answer yes or no only." - ) - prompt = template.format(question, answer, response) - elif task == "temporal-reasoning": - template = ( - "I will give you a question, a correct answer, and a response from a model. Please answer yes " - "if the response contains the correct answer. Otherwise, answer no. If the response is equiva" - "lent to the correct answer or contains all the intermediate steps to get the correct answer" - ", you should also answer yes. If the response only contains a subset of the information requi" - "red by the answer, answer no. In addition, do not penalize off-by-one errors for the number" - " of days. If the question asks for the number of days/weeks/months, etc., and the model makes" - " off-by-one errors (e.g., predicting 19 days when the answer is 18), the model's response is" - " still correct. \n\nQuestion: {}\n\nCorrect Answer: {}\n\nModel Response: {}\n\nIs the mode" - "l response correct? Answer yes or no only." - ) - prompt = template.format(question, answer, response) - elif task == "knowledge-update": - template = ( - "I will give you a question, a correct answer, and a response from a model. Please answer ye" - "s if the response contains the correct answer. Otherwise, answer no. If the response contai" - "ns some previous information along with an updated answer, the response should be consider" - "ed as correct as long as the updated answer is the required answer.\n\nQuestion: {}\n\nCo" - "rrect Answer: {}\n\nModel Response: {}\n\nIs the model response correct? Answer yes or no" - " only." - ) - prompt = template.format(question, answer, response) - elif task == "single-session-preference": - template = ( - "I will give you a question, a rubric for desired personalized response, and a response fro" - "m a model. Please answer yes if the response satisfies the desired response. Otherwise, ans" - "wer no. The model does not need to reflect all the points in the rubric. The response is corr" - "ect as long as it recalls and utilizes the user's personal information correctly.\n\nQues" - "tion: {}\n\nRubric: {}\n\nModel Response: {}\n\nIs the model response correct? Answer yes" - " or no only." - ) - prompt = template.format(question, answer, response) - else: - # Default template - template = ( - "I will give you a question, a correct answer, and a response from a model. Please answer yes" - " if the response contains the correct answer. Otherwise, answer no. If the response is equival" - "ent to the correct answer or contains all the intermediate steps to get the correct ans" - "wer, you should also answer yes. If the response only contains a subset of the information" - " required by the answer, answer no. \n\nQuestion: {}\n\nCorrect Answer: {}\n\nModel Response:" - " {}\n\nIs the model response correct? Answer yes or no only." - ) - prompt = template.format(question, answer, response) - - else: - template = ( - "I will give you an unanswerable question, an explanation, and a response from a mode" - "l. Please answer yes if the model correctly identifies the question as unanswerable. The model " - "could say that the information is incomplete, or some other information is given but the asked " - "information is not.\n\nQuestion: {}\n\nExplanation: {}\n\nModel Response: {}\n\nDoes the model " - "correctly identify the question as unanswerable? Answer yes or no only." - ) - prompt = template.format(question, answer, response) - return prompt - - -# ==================== Utilities ==================== - - -class DataLoader: - """Handles loading and parsing of LongMemEval data.""" - - @staticmethod - def load_json(file_path: str) -> list[dict]: - """Load all entries from a JSON file.""" - with open(file_path, "r", encoding="utf-8") as f: - return json.load(f) - - @staticmethod - def filter_by_type(data: list[dict], samples_per_type: int = -1) -> list[tuple[int, dict]]: - """Filter data by question type with specified number of samples per type. - - Args: - data: List of question entries - samples_per_type: Number of samples per type, -1 for all - - Returns: - List of tuples (original_index, entry) for selected samples - """ - if samples_per_type == -1: - # Return all with original indices - return list(enumerate(data)) - - # Group by question type - type_groups: dict[str, list[tuple[int, dict]]] = {} - for i, entry in enumerate(data): - qtype = entry.get("question_type", "unknown") - if qtype not in type_groups: - type_groups[qtype] = [] - type_groups[qtype].append((i, entry)) - - # Select samples from each type - selected = [] - for qtype, entries in type_groups.items(): - count = min(samples_per_type, len(entries)) - selected.extend(entries[:count]) - logger.info(f" {qtype}: selected {count}/{len(entries)} samples") - - # Sort by original index to maintain order - selected.sort(key=lambda x: x[0]) - return selected - - @staticmethod - def convert_session_to_messages(session: list[dict], session_date: str) -> list[dict]: - """Convert LongMemEval session to ReMe message format. - - Args: - session: List of messages, each containing role, content, has_answer - session_date: Session date in format '2023/04/10 (Mon) 17:50' - - Returns: - List of messages with time_created field (user messages only) - """ - messages = [] - - # Parse session date as base time - try: - # Format: "2023/04/10 (Mon) 17:50" - date_part = session_date.split(" (")[0] - time_part = session_date.split(") ")[1] if ") " in session_date else "00:00" - base_time = datetime.strptime(f"{date_part} {time_part}", "%Y/%m/%d %H:%M") - except Exception: - base_time = datetime.now() - - for i, msg in enumerate(session): - # Add 1 minute per message - msg_time = base_time + timedelta(minutes=i) - - messages.append( - { - "role": msg["role"], - "content": msg["content"], - "time_created": msg_time.replace(tzinfo=timezone.utc).strftime("%Y-%m-%d %H:%M:%S"), - }, - ) - - return messages - - -class FileManager: - """Manages file I/O operations.""" - - def __init__(self, base_dir: str): - self.base_dir = Path(base_dir) - self.base_dir.mkdir(parents=True, exist_ok=True) - - def save_question_result(self, idx: int, question_id: str, data: dict): - """Save result for a single question.""" - file_path = self.base_dir / f"question_{idx:04d}_{question_id}.json" - with open(file_path, "w", encoding="utf-8") as f: - json.dump(data, f, indent=4, ensure_ascii=False) - logger.info(f"✅ Saved question result to {file_path}") - - def load_question_result(self, idx: int, question_id: str) -> Optional[dict]: - """Load result for a single question if exists.""" - file_path = self.base_dir / f"question_{idx:04d}_{question_id}.json" - if not file_path.exists(): - return None - with open(file_path, "r", encoding="utf-8") as f: - return json.load(f) - - def save_summary(self, results: list[dict]): - """Save summary of all results.""" - file_path = self.base_dir / "summary.json" - with open(file_path, "w", encoding="utf-8") as f: - json.dump(results, f, indent=4, ensure_ascii=False) - logger.info(f"✅ Saved summary to {file_path}") - - -# ==================== Evaluation Functions ==================== - - -async def answer_question_with_memories( - reme: ReMe, - question: str, - memories: str, - user_id: str = None, - model_name: str = "qwen-max", -): - """ - Answer a question using retrieved memories with PROMPT_MEMZERO_JSON template. - - Args: - reme: ReMe instance with default_llm and prompt_handler - question: The question to answer - memories: The retrieved memories (formatted as context) - user_id: Optional user ID for context formatting - model_name: Model name to use for LLM request - - Returns: - dict with 'reasoning' and 'answer' fields - """ - # Format context with memories - if user_id: - context = reme.prompt_handler.prompt_format( - "TEMPLATE_MEMOS", - user_id=user_id, - memories=memories, - ) - else: - context = f"Memories:\n{memories}" - - # Use PROMPT_MEMZERO_JSON template for structured JSON response - prompt = reme.prompt_handler.prompt_format( - "PROMPT_MEMZERO_JSON", - context=context, - question=question, - ) - - result = await reme.default_llm.simple_request_for_json( - prompt=prompt, - model_name=model_name, - ) - - return result - - -# ==================== Memory Operations ==================== - - -class MemoryProcessor: - """Handles ReMe memory operations.""" - - def __init__( - self, - reme: ReMe, - reme_model_name: str = "qwen-flash", - retrieve_model_name: str = "qwen-max", - eval_model_name: str = "qwen-max", - algo_version: str = "v1", - enable_thinking_params: bool = False, - ): - self.reme = reme - self.reme_model_name = reme_model_name - self.retrieve_model_name = retrieve_model_name - self.eval_model_name = eval_model_name - self.algo_version = algo_version - self.enable_thinking_params = enable_thinking_params - - async def add_memories( - self, - user_id: str, - messages: list[dict], - batch_size: int = 10000, - ) -> tuple[list[dict], list, float]: - """ - Add memories in batches using ReMe and return extracted memory contents. - - Returns: - tuple: (extracted_memories, agent_messages, total_duration_ms) - """ - extracted_memories = [] - summary_messages = [] - total_duration_ms = 0 - - for i in range(0, len(messages), batch_size): - batch = messages[i : i + batch_size] - start = time.time() - - # Use new summary API - result = await self.reme.summarize_memory( - messages=batch, - user_name=user_id, - version=self.algo_version, - return_dict=True, - enable_time_filter=True, - enable_thinking_params=self.enable_thinking_params, - ) - - duration_ms = (time.time() - start) * 1000 - total_duration_ms += duration_ms - - extracted_memories.extend([m.model_dump(exclude_none=True) for m in result["answer"]]) - summary_messages.extend([m.simple_dump(enable_argument_dict=True) for m in result["messages"]]) - - return extracted_memories, summary_messages, total_duration_ms - - async def search_memory( - self, - query: str, - user_id: str, - top_k: int = 20, - ) -> tuple[dict, list, float]: - """ - Search memory using ReMe and return structured answer with reasoning. - - Returns: - tuple: (answer_dict, agent_messages, duration_ms) - answer_dict contains: {"reasoning": str, "answer": str, "memories": str} - """ - start = time.time() - - # Retrieve memories from ReMe using new API - result = await self.reme.retrieve_memory( - llm_config_name=self.retrieve_model_name, - query=query, - retrieve_top_k=top_k, - user_name=user_id, - version=self.algo_version, - return_dict=True, - enable_time_filter=True, - enable_thinking_params=self.enable_thinking_params, - ) - - # Extract memories from response - memories = result["answer"] - agent_messages = [x.simple_dump(enable_argument_dict=True) for x in result["messages"]] - retrieved_nodes = [x.model_dump(exclude_none=True) for x in result["retrieved_nodes"]] - - # Use LLM to generate structured answer from memories - answer_result = await answer_question_with_memories( - reme=self.reme, - question=query, - memories=memories, - user_id=user_id, - model_name=self.eval_model_name, - ) - - # Add original memories to the result - answer_result["memories"] = memories - answer_result["retrieved_nodes"] = retrieved_nodes - - duration_ms = (time.time() - start) * 1000 - return answer_result, agent_messages, duration_ms - - -# ==================== Answer Judge ==================== - - -class LongMemEvalJudge: - """LongMemEval answer judge using LLM.""" - - def __init__(self, reme: ReMe, model: str = "qwen3-max"): - self.reme = reme - self.model = model - - async def judge_answer( - self, - question_type: str, - question: str, - answer: str, - response: str, - abstention: bool = False, - ) -> dict: - """ - Judge if the model's response is correct. - - Returns: - dict with is_correct, llm_response, and judge_prompt - """ - prompt = get_anscheck_prompt(question_type, question, answer, response, abstention) - - try: - llm_response = await self.reme.get_llm("default").simple_request( - prompt=prompt, - model_name=self.model, - ) - llm_response_lower = llm_response.strip().lower() - is_correct = llm_response_lower.startswith("yes") - - return { - "is_correct": is_correct, - "llm_response": llm_response, - "judge_prompt": prompt, - } - except Exception as e: - return { - "is_correct": None, - "error": str(e), - "judge_prompt": prompt, - } - - -# ==================== Metrics ==================== - - -class MetricsAggregator: - """Aggregates evaluation metrics for LongMemEval.""" - - @staticmethod - def compute_metrics(results: list[dict]) -> dict[str, Any]: - """Compute overall and per-type metrics.""" - total = len(results) - correct = sum(1 for r in results if r.get("judgment", {}).get("is_correct") is True) - incorrect = sum(1 for r in results if r.get("judgment", {}).get("is_correct") is False) - error = total - correct - incorrect - - metrics = { - "total": total, - "correct": correct, - "incorrect": incorrect, - "error": error, - "accuracy": correct / total if total > 0 else 0, - "accuracy_valid": correct / (correct + incorrect) if (correct + incorrect) > 0 else 0, - } - - # Per question type statistics - type_stats = {} - for r in results: - qtype = r.get("question_type", "unknown") - if qtype not in type_stats: - type_stats[qtype] = {"total": 0, "correct": 0, "incorrect": 0} - type_stats[qtype]["total"] += 1 - if r.get("judgment", {}).get("is_correct") is True: - type_stats[qtype]["correct"] += 1 - elif r.get("judgment", {}).get("is_correct") is False: - type_stats[qtype]["incorrect"] += 1 - - metrics["by_question_type"] = { - qtype: { - **stats, - "accuracy": stats["correct"] / stats["total"] if stats["total"] > 0 else 0, - "accuracy_valid": ( - stats["correct"] / (stats["correct"] + stats["incorrect"]) - if (stats["correct"] + stats["incorrect"]) > 0 - else 0 - ), - } - for qtype, stats in type_stats.items() - } - - return metrics - - @staticmethod - def compute_timing_stats(results: list[dict]) -> dict[str, Any]: - """Compute timing statistics.""" - summary_times = [] - retrieve_times = [] - - for r in results: - summary_ms = r.get("summary_duration_ms", 0) - retrieve_ms = r.get("retrieve_duration_ms", 0) - - if summary_ms > 0: - summary_times.append(summary_ms) - if retrieve_ms > 0: - retrieve_times.append(retrieve_ms) - - def compute_stats(times: list[float]) -> dict: - if not times: - return {"count": 0, "total_ms": 0, "avg_ms": 0, "min_ms": 0, "max_ms": 0} - - return { - "count": len(times), - "total_ms": sum(times), - "total_min": sum(times) / 1000 / 60, - "avg_ms": sum(times) / len(times), - "min_ms": min(times), - "max_ms": max(times), - } - - return { - "summary": compute_stats(summary_times), - "retrieve": compute_stats(retrieve_times), - "total_time_min": (sum(summary_times) + sum(retrieve_times)) / 1000 / 60, - } - - -# ==================== Main Pipeline ==================== - - -class LongMemEvalEvaluator: - """Main evaluator for LongMemEval benchmark using ReMe.""" - - def __init__(self, config: EvalConfig): - self.config = config - self.file_manager = FileManager(config.output_dir) - self.data_loader = DataLoader() - - # Store LLM configs for creating ReMe instances per question - self._llm_configs = { - "qwen-plus-t": { - "backend": "openai", - "model_name": "qwen-plus", - "extra_body": { - "enable_thinking": True, - }, - }, - "qwen-max-t": { - "backend": "openai", - "model_name": "qwen3-max", - "extra_body": { - "enable_thinking": True, - }, - }, - "gpt-4o-mini": { - "backend": "openai", - "model_name": "gpt-4o-mini-2024-07-18", - }, - "gpt-4o-mini-2024-07-18": { - "backend": "openai", - "model_name": "gpt-4o-mini-2024-07-18", - }, - "qwen-flash": { - "backend": "openai", - "model_name": "qwen-flash", - }, - "qwen-max": { - "backend": "openai", - "model_name": "qwen3-max", - }, - } - - # Load evaluation prompts path - self._prompts_yaml_path = Path(__file__).parent / "eval_reme.yaml" - - def _create_reme_for_question(self, question_id: str) -> ReMe: - """Create a ReMe instance for a specific question with isolated collection. - - Args: - question_id: The question ID to use as collection name - - Returns: - ReMe instance with isolated vector store collection - """ - collection_name = f"longmemeval_{question_id}" - reme = ReMe( - default_llm_config={ - "model_name": self.config.reme_model_name, - }, - default_vector_store_config={ - "collection_name": collection_name, - }, - llms=self._llm_configs, - ) - - # Load evaluation prompts - reme.prompt_handler.load_prompt_by_file(self._prompts_yaml_path) - - return reme - - async def __aenter__(self): - """Async context manager entry.""" - return self - - async def __aexit__(self, exc_type, exc_val, exc_tb): - """Async context manager exit with cleanup.""" - return False - - async def process_question_entry(self, entry: dict, idx: int) -> dict: - """Process a single question entry. - - Each question gets its own ReMe instance with isolated vector store collection. - - Args: - entry: A question entry from LongMemEval dataset - idx: Index of the question - - Returns: - Result dictionary - """ - question_id = entry["question_id"] - question = entry["question"] - answer = entry["answer"] - question_type = entry["question_type"] - question_date = entry.get("question_date", "") - haystack_dates = entry["haystack_dates"] - haystack_session_ids = entry["haystack_session_ids"] - haystack_sessions = entry["haystack_sessions"] - - # Use "User" as user_name, question_id is stored in collection_name - user_name = "User" - - logger.info(f"\n{'=' * 60}") - logger.info(f"Question ID: {question_id}") - logger.info(f"Question Type: {question_type}") - logger.info(f"Question: {question}") - logger.info(f"Question_date: {question_date}") - logger.info(f"Answer: {answer}") - logger.info(f"Number of sessions: {len(haystack_sessions)}") - logger.info(f"{'=' * 60}") - - # Create isolated ReMe instance for this question - reme = self._create_reme_for_question(question_id) - await reme.start() - - try: - # Create memory processor and judge for this ReMe instance - memory_processor = MemoryProcessor( - reme, - self.config.reme_model_name, - self.config.retrieve_model_name, - self.config.eval_model_name, - self.config.algo_version, - self.config.enable_thinking_params, - ) - judge = LongMemEvalJudge(reme, self.config.eval_model_name) - - # Clear existing vector store data for this collection - await reme.default_vector_store.delete_all() - - # Step 2: Process all haystack sessions to build memory - all_extracted_memories = [] - all_agent_messages = [] - total_summary_duration_ms = 0 - - for session_idx, (session, session_date, session_id) in enumerate( - zip(haystack_sessions, haystack_dates, haystack_session_ids), - ): - logger.info(f" Processing session {session_idx + 1}/{len(haystack_sessions)}: {session_id}") - - # Convert session to messages - messages = self.data_loader.convert_session_to_messages(session, session_date) - - if not messages: - continue - - # Add memories using "User" as user_name - extracted_memories, agent_messages, duration_ms = await memory_processor.add_memories( - user_id=user_name, - messages=messages, - batch_size=self.config.batch_size, - ) - - all_extracted_memories.extend(extracted_memories) - all_agent_messages.extend(agent_messages) - total_summary_duration_ms += duration_ms - - # Step 3: Search memory and answer question - logger.info(" Answering question using ReMe...") - answer_dict, retrieve_messages, retrieve_duration_ms = await memory_processor.search_memory( - query=f"[Question_date: {question_date} | Question_type: {question_type}] " + question, - user_id=user_name, - top_k=self.config.top_k, - ) - - # Extract answer and reasoning from the structured response - model_response = answer_dict.get("answer", "") - model_reasoning = answer_dict.get("reasoning", "") - retrieved_memories = answer_dict.get("memories", "") - retrieved_nodes = answer_dict.get("retrieved_nodes", []) - - # Step 4: Judge answer correctness - logger.info(" Judging answer correctness...") - judgment = await judge.judge_answer( - question_type=question_type, - question=question, - answer=answer, - response=model_response, - ) - - is_correct = judgment.get("is_correct") - logger.info( - f" → Answer judgment: {'Correct' if is_correct else 'Incorrect' if is_correct is False else 'Error'}", - ) - - result = { - "question_id": question_id, - "question_type": question_type, - "question": question, - "answer": answer, - "question_date": question_date, - "haystack_dates": haystack_dates, - "haystack_session_ids": haystack_session_ids, - "num_sessions": len(haystack_sessions), - "model_response": model_response, - "model_reasoning": model_reasoning, - "retrieved_memories": retrieved_memories, - "retrieved_nodes": retrieved_nodes, - "judgment": judgment, - "extracted_memories": all_extracted_memories, - "summary_duration_ms": total_summary_duration_ms, - "retrieve_duration_ms": retrieve_duration_ms, - "summary_messages": all_agent_messages, - "retrieve_messages": retrieve_messages, - } - - # Save individual result - self.file_manager.save_question_result(idx, question_id, result) - - logger.info(f" Question {question_id} - Completed") - - return result - - finally: - # Always close the ReMe instance - await reme.close() - - async def run_evaluation(self): - """Run the complete evaluation pipeline with parallel processing.""" - start_time = time.time() - - # Load dataset - logger.info(f"Loading dataset from: {self.config.data_path}") - all_data = self.data_loader.load_json(self.config.data_path) - logger.info(f"Total questions in dataset: {len(all_data)}") - - # Filter by question type - logger.info(f"Filtering by type (samples_per_type={self.config.samples_per_type}):") - filtered_data = self.data_loader.filter_by_type(all_data, self.config.samples_per_type) - logger.info(f"Selected {len(filtered_data)} questions after filtering") - - # Apply start_index and end_index on filtered data - end_index = self.config.end_index or len(filtered_data) - start_index = self.config.start_index - end_index = min(end_index, len(filtered_data)) - - # Get the slice we want to process - data_to_process = filtered_data[start_index:end_index] - total_questions = len(data_to_process) - - logger.info(f"Processing {total_questions} questions (index {start_index} to {end_index - 1})") - - print("\n" + "=" * 80) - print("LONGMEMEVAL EVALUATION - REME") - print(f"Samples per type: {self.config.samples_per_type} (-1 = all)") - print(f"Questions to process: {total_questions} | Top-K: {self.config.top_k}") - print(f"Max Concurrency: {self.config.max_concurrency}") - print( - f"Summary Model: {self.config.reme_model_name} | Retrieve Model: {self.config.retrieve_model_name} " - f"| Eval Model: {self.config.eval_model_name}", - ) - print(f"Algo Version: {self.config.algo_version}") - print("=" * 80 + "\n") - - # Use semaphore to control concurrency - semaphore = asyncio.Semaphore(self.config.max_concurrency) - - async def process_with_semaphore(idx: int, original_idx: int, entry: dict) -> Optional[dict]: - """Process a question with semaphore for concurrency control.""" - async with semaphore: - question_id = entry["question_id"] - - # Check cache first (use original index for cache file naming) - cached_result = self.file_manager.load_question_result(original_idx, question_id) - if cached_result: - print(f"⚡ [{idx}/{total_questions}] Skipping question {original_idx} (cached)") - return cached_result - - print(f"\n{'#' * 60}") - print(f"### [{idx}/{total_questions}] Processing Question {original_idx} ###") - print(f"{'#' * 60}") - - try: - result = await self.process_question_entry(entry, original_idx) - print(f"✅ [{idx}/{total_questions}] Completed question {original_idx}") - return result - except Exception as e: - logger.error(f"❌ Error processing question {original_idx}: {e}") - import traceback - - traceback.print_exc() - return { - "question_id": question_id, - "error": str(e), - "question_type": entry.get("question_type", "unknown"), - "question": entry.get("question", ""), - "answer": entry.get("answer", ""), - "judgment": {"is_correct": None, "error": str(e)}, - } - - # Create all tasks from filtered data (each item is a tuple of (original_idx, entry)) - tasks = [ - process_with_semaphore(idx + 1, original_idx, entry) - for idx, (original_idx, entry) in enumerate(data_to_process) - ] - - # Execute in parallel with controlled concurrency - all_results = await asyncio.gather(*tasks, return_exceptions=False) - - # Filter out None results if any - all_results = [r for r in all_results if r is not None] - - # Save summary - self.file_manager.save_summary(all_results) - - elapsed = time.time() - start_time - print(f"\n✅ Processing completed in {elapsed:.2f}s") - if total_questions > 0: - print(f" Average time per question: {elapsed / total_questions:.2f}s") - - # Compute and report metrics - self._report_metrics(all_results) - - return all_results - - def _report_metrics(self, results: list[dict]): - """Report evaluation metrics.""" - metrics = MetricsAggregator.compute_metrics(results) - timing_stats = MetricsAggregator.compute_timing_stats(results) - - print("\n" + "=" * 80) - print("EVALUATION SUMMARY - LONGMEMEVAL - REME") - print("=" * 80 + "\n") - - print("📊 Overall Results:") - print(f" ✅ Correct: {metrics['correct']}/{metrics['total']} ({100 * metrics['accuracy']:.2f}%)") - print( - f" ❌ Incorrect: {metrics['incorrect']}/{metrics['total']} " - f"({100 * metrics['incorrect'] / metrics['total'] if metrics['total'] > 0 else 0:.2f}%)", - ) - if metrics["error"] > 0: - print( - f" ⚠️ Error: {metrics['error']}/{metrics['total']} ({100 * metrics['error'] / metrics['total']:.2f}%)", - ) - print(f" Accuracy (valid): {100 * metrics['accuracy_valid']:.2f}%") - - print("\n📊 Accuracy by Question Type:") - print("-" * 60) - print(f"{'Question Type':<30} {'Correct':<10} {'Total':<10} {'Accuracy':<10}") - print("-" * 60) - for qtype in sorted(metrics["by_question_type"].keys()): - stats = metrics["by_question_type"][qtype] - print(f"{qtype:<30} {stats['correct']:<10} {stats['total']:<10} {100 * stats['accuracy']:.2f}%") - print("-" * 60) - - print("\n⏱️ Timing Statistics:") - summary = timing_stats["summary"] - retrieve = timing_stats["retrieve"] - print(" Memory Summarization:") - print(f" Total Time: {summary['total_ms']:.2f} min") - print(f" Avg per Q: {summary['avg_ms']:.0f} ms") - print(" Memory Retrieval:") - print(f" Total Time: {retrieve['total_ms']:.2f} min") - print(f" Avg per Q: {retrieve['avg_ms']:.0f} ms") - print(f" Total Time: {timing_stats['total_time_min']:.2f} min") - - # Save metrics - final_results = { - "accuracy": metrics, - "timing": timing_stats, - } - metrics_file = self.file_manager.base_dir / "eval_statistics.json" - with open(metrics_file, "w", encoding="utf-8") as f: - json.dump(final_results, f, indent=4, ensure_ascii=False) - print(f"\n📁 Statistics saved to: {metrics_file}") - - print("\n" + "=" * 80) - - -# ==================== Entry Point ==================== - - -async def main_async( - data_path: str, - top_k: int = 20, - start_index: int = 0, - end_index: Optional[int] = None, - max_concurrency: int = 1, - batch_size: int = 30, - output_dir: str = "bench_results/longmemeval_reme", - reme_model_name: str = "qwen-flash", - retrieve_model_name: str = "qwen-max", - eval_model_name: str = "qwen-max", - algo_version: str = "v1", - samples_per_type: int = -1, - enable_thinking_params: bool = False, -): - """Main async entry point for LongMemEval evaluation with proper resource cleanup.""" - config = EvalConfig( - data_path=data_path, - top_k=top_k, - start_index=start_index, - end_index=end_index, - max_concurrency=max_concurrency, - batch_size=batch_size, - output_dir=output_dir, - reme_model_name=reme_model_name, - retrieve_model_name=retrieve_model_name, - eval_model_name=eval_model_name, - algo_version=algo_version, - samples_per_type=samples_per_type, - enable_thinking_params=enable_thinking_params, - ) - - # Use async context manager for automatic cleanup - async with LongMemEvalEvaluator(config) as evaluator: - await evaluator.run_evaluation() - - -def main( - data_path: str, - top_k: int = 20, - start_index: int = 0, - end_index: Optional[int] = None, - max_concurrency: int = 1, - batch_size: int = 30, - output_dir: str = "bench_results/longmemeval_reme", - reme_model_name: str = "qwen-flash", - retrieve_model_name: str = "qwen-max", - eval_model_name: str = "qwen-max", - algo_version: str = "v1", - samples_per_type: int = -1, - enable_thinking_params: bool = False, -): - """Main entry point for LongMemEval evaluation.""" - asyncio.run( - main_async( - data_path=data_path, - top_k=top_k, - start_index=start_index, - end_index=end_index, - max_concurrency=max_concurrency, - batch_size=batch_size, - output_dir=output_dir, - reme_model_name=reme_model_name, - retrieve_model_name=retrieve_model_name, - eval_model_name=eval_model_name, - algo_version=algo_version, - samples_per_type=samples_per_type, - enable_thinking_params=enable_thinking_params, - ), - ) - - -if __name__ == "__main__": - import argparse - - parser = argparse.ArgumentParser( - description="Evaluate ReMe on LongMemEval benchmark", - ) - parser.add_argument( - "--data_path", - type=str, - required=True, - help="Path to LongMemEval JSON file", - ) - parser.add_argument( - "--top_k", - type=int, - default=10, - help="Number of memories to retrieve (default: 20)", - ) - parser.add_argument( - "--start_index", - type=int, - default=0, - help="Start index for processing questions (default: 0)", - ) - parser.add_argument( - "--end_index", - type=int, - default=None, - help="End index for processing questions (default: None, process all)", - ) - parser.add_argument( - "--max_concurrency", - type=int, - default=4, - help="Maximum concurrent question processing (default: 1)", - ) - parser.add_argument( - "--batch_size", - type=int, - default=30, - help="Batch size for memory summary processing (default: 30)", - ) - parser.add_argument( - "--output_dir", - type=str, - default="bench_results/longmemeval_reme", - help="Output directory for results", - ) - parser.add_argument( - "--reme_model_name", - type=str, - default="qwen-flash", - help="Model name for ReMe summary operations (default: qwen-flash)", - ) - parser.add_argument( - "--retrieve_model_name", - type=str, - default="qwen-max", - help="Model name for memory retrieval (default: qwen-max)", - ) - parser.add_argument( - "--eval_model_name", - type=str, - default="qwen-flash", - help="Model name for evaluation/judgment (default: qwen-max)", - ) - parser.add_argument( - "--algo_version", - type=str, - default="default", - help="Algorithm version for summary and retrieval (default: v1)", - ) - parser.add_argument( - "--samples_per_type", - type=int, - default=1, - help="Number of samples per question type, -1 for all (default: -1)", - ) - parser.add_argument( - "--enable_thinking_params", - action="store_true", - default=False, - help="Enable thinking parameters for summary and retrieval (default: False)", - ) - - args = parser.parse_args() - print(f"args={args}!") - - main( - data_path=args.data_path, - top_k=args.top_k, - start_index=args.start_index, - end_index=args.end_index, - max_concurrency=args.max_concurrency, - batch_size=args.batch_size, - output_dir=args.output_dir, - reme_model_name=args.reme_model_name, - retrieve_model_name=args.retrieve_model_name, - eval_model_name=args.eval_model_name, - algo_version=args.algo_version, - samples_per_type=args.samples_per_type, - enable_thinking_params=args.enable_thinking_params, - ) diff --git a/benchmark/longmemeval/eval_longmemeval_reme_retrieve.py b/benchmark/longmemeval/eval_longmemeval_reme_retrieve.py deleted file mode 100644 index 2e528cbd..00000000 --- a/benchmark/longmemeval/eval_longmemeval_reme_retrieve.py +++ /dev/null @@ -1,921 +0,0 @@ -""" -LongMemEval Benchmark Evaluator for ReMe - Retrieve Only - -A simplified evaluation pipeline that only runs the retrieve and judge phases: -1. Loads LongMemEval benchmark data -2. Skips memory summarization (assumes memories are already in vector store) -3. Uses questions to query memory and generate answers -4. Uses LLM to judge answer correctness -5. Generates comprehensive metrics - -This is useful for debugging/tuning the retrieve phase without re-running summary. - -Usage: - python benchmark/longmemeval/eval_longmemeval_reme_retrieve.py \ - --data_path dataset/longmemeval/longmemeval_s_cleaned.json \ - --top_k 20 --start_index 0 --end_index 10 -""" - -import asyncio -import json -import time -from dataclasses import dataclass -from pathlib import Path -from typing import Any, Optional - -from loguru import logger - -from reme.reme import ReMe - - -# ==================== Configuration ==================== - - -@dataclass -class RetrieveEvalConfig: - """Evaluation configuration parameters for retrieve-only mode.""" - - data_path: str - top_k: int = 10 - start_index: int = 0 - end_index: Optional[int] = None - max_concurrency: int = 1 - output_dir: str = "cache/bench_results/longmemeval_reme_retrieve" - reme_model_name: str = "qwen-flash" - eval_model_name: str = "qwen3-max" - algo_version: str = "v1" - samples_per_type: int = -1 # Number of samples per question type, -1 for all - enable_thinking_params: bool = False - # Optional: path to previous summary results to reload memories - summary_results_dir: Optional[str] = None - - -# ==================== Answer Judge Prompts ==================== - - -def get_anscheck_prompt(task: str, question: str, answer: str, response: str, abstention: bool = False) -> str: - """Generate the answer checking prompt based on question type. - - Args: - task: Question type, e.g. 'single-session-user', 'multi-session', 'temporal-reasoning' - question: The question content - answer: The reference answer - response: The model's response - abstention: Whether this is an unanswerable question - - Returns: - Prompt for judging answer correctness - """ - if not abstention: - if task in ["single-session-user", "single-session-assistant", "multi-session"]: - template = ( - "I will give you a question, a correct answer, and a response from a model. Please answer yes i" - "f the response contains the correct answer. Otherwise, answer no. If the response is equival" - "ent to the correct answer or contains all the intermediate steps to get the correct answer, " - "you should also answer yes. If the response only contains a subset of the information required" - " by the answer, answer no. \n\nQuestion: {}\n\nCorrect Answer: {}\n\nModel Response: {}\n\nIs" - " the model response correct? Answer yes or no only." - ) - prompt = template.format(question, answer, response) - elif task == "temporal-reasoning": - template = ( - "I will give you a question, a correct answer, and a response from a model. Please answer yes" - " if the response contains the correct answer. Otherwise, answer no. If the response is equiv" - "alent to the correct answer or contains all the intermediate steps to get the correct answer" - ", you should also answer yes. If the response only contains a subset of the information requ" - "ired by the answer, answer no. In addition, do not penalize off-by-one errors for the numbe" - "r of days. If the question asks for the number of days/weeks/months, etc., and the model ma" - "kes off-by-one errors (e.g., predicting 19 days when the answer is 18), the model's respon" - "se is still correct. \n\nQuestion: {}\n\nCorrect Answer: {}\n\nModel Response: {}\n\nIs th" - "e model response correct? Answer yes or no only." - ) - prompt = template.format(question, answer, response) - elif task == "knowledge-update": - template = ( - "I will give you a question, a correct answer, and a response from a model. Please answer yes " - "if the response contains the correct answer. Otherwise, answer no. If the response contains " - "some previous information along with an updated answer, the response should be considered " - "as correct as long as the updated answer is the required answer.\n\nQuestion: {}\n\nCorrec" - "t Answer: {}\n\nModel Response: {}\n\nIs the model response correct? Answer yes or no only." - ) - prompt = template.format(question, answer, response) - elif task == "single-session-preference": - template = ( - "I will give you a question, a rubric for desired personalized response, and a response from a" - " model. Please answer yes if the response satisfies the desired response. Otherwise, answer" - " no. The model does not need to reflect all the points in the rubric. The response is corr" - "ect as long as it recalls and utilizes the user's personal information correctly.\n\nQuest" - "ion: {}\n\nRubric: {}\n\nModel Response: {}\n\nIs the model response correct? Answer yes o" - "r no only." - ) - prompt = template.format(question, answer, response) - else: - # Default template - template = ( - "I will give you a question, a correct answer, and a response from a model. Please answer y" - "es if the response contains the correct answer. Otherwise, answer no. If the response is " - "equivalent to the correct answer or contains all the intermediate steps to get the correc" - "t answer, you should also answer yes. If the response only contains a subset of the infor" - "mation required by the answer, answer no. \n\nQuestion: {}\n\nCorrect Answer: {}\n\nModel" - " Response: {}\n\nIs the model response correct? Answer yes or no only." - ) - prompt = template.format(question, answer, response) - - else: - template = ( - "I will give you an unanswerable question, an explanation, and a response from a model. Please " - "answer yes if the model correctly identifies the question as unanswerable. The model could say " - "that the information is incomplete, or some other information is given but the asked informati" - "on is not.\n\nQuestion: {}\n\nExplanation: {}\n\nModel Response: {}\n\nDoes the model correct" - "ly identify the question as unanswerable? Answer yes or no only." - ) - prompt = template.format(question, answer, response) - return prompt - - -# ==================== Utilities ==================== - - -class DataLoader: - """Handles loading and parsing of LongMemEval data.""" - - @staticmethod - def load_json(file_path: str) -> list[dict]: - """Load all entries from a JSON file.""" - with open(file_path, "r", encoding="utf-8") as f: - return json.load(f) - - @staticmethod - def filter_by_type(data: list[dict], samples_per_type: int = -1) -> list[tuple[int, dict]]: - """Filter data by question type with specified number of samples per type. - - Args: - data: List of question entries - samples_per_type: Number of samples per type, -1 for all - - Returns: - List of tuples (original_index, entry) for selected samples - """ - if samples_per_type == -1: - # Return all with original indices - return list(enumerate(data)) - - # Group by question type - type_groups: dict[str, list[tuple[int, dict]]] = {} - for i, entry in enumerate(data): - qtype = entry.get("question_type", "unknown") - if qtype not in type_groups: - type_groups[qtype] = [] - type_groups[qtype].append((i, entry)) - - # Select samples from each type - selected = [] - for qtype, entries in type_groups.items(): - count = min(samples_per_type, len(entries)) - selected.extend(entries[:count]) - logger.info(f" {qtype}: selected {count}/{len(entries)} samples") - - # Sort by original index to maintain order - selected.sort(key=lambda x: x[0]) - return selected - - -class FileManager: - """Manages file I/O operations.""" - - def __init__(self, base_dir: str): - self.base_dir = Path(base_dir) - self.base_dir.mkdir(parents=True, exist_ok=True) - - def save_question_result(self, idx: int, question_id: str, data: dict): - """Save result for a single question.""" - file_path = self.base_dir / f"question_{idx:04d}_{question_id}.json" - with open(file_path, "w", encoding="utf-8") as f: - json.dump(data, f, indent=4, ensure_ascii=False) - logger.info(f"✅ Saved question result to {file_path}") - - def load_question_result(self, idx: int, question_id: str) -> Optional[dict]: - """Load result for a single question if exists.""" - file_path = self.base_dir / f"question_{idx:04d}_{question_id}.json" - if not file_path.exists(): - return None - with open(file_path, "r", encoding="utf-8") as f: - return json.load(f) - - def save_summary(self, results: list[dict]): - """Save summary of all results.""" - file_path = self.base_dir / "summary.json" - with open(file_path, "w", encoding="utf-8") as f: - json.dump(results, f, indent=4, ensure_ascii=False) - logger.info(f"✅ Saved summary to {file_path}") - - -# ==================== Evaluation Functions ==================== - - -async def answer_question_with_memories( - reme: ReMe, - question: str, - memories: str, - user_id: str = None, - model_name: str = "qwen3-max", -): - """ - Answer a question using retrieved memories with PROMPT_MEMZERO_JSON template. - - Args: - reme: ReMe instance with default_llm and prompt_handler - question: The question to answer - memories: The retrieved memories (formatted as context) - user_id: Optional user ID for context formatting - model_name: Model name to use for LLM request - - Returns: - dict with 'reasoning' and 'answer' fields - """ - # Format context with memories - if user_id: - context = reme.prompt_handler.prompt_format( - "TEMPLATE_MEMOS", - user_id=user_id, - memories=memories, - ) - else: - context = f"Memories:\n{memories}" - - # Use PROMPT_MEMZERO_JSON template for structured JSON response - prompt = reme.prompt_handler.prompt_format( - "PROMPT_MEMZERO_JSON", - context=context, - question=question, - ) - - result = await reme.get_llm(name=model_name).simple_request_for_json( - prompt=prompt, - model_name=None, - ) - - return result - - -# ==================== Memory Operations ==================== - - -class RetrieveProcessor: - """Handles ReMe memory retrieve operations only.""" - - def __init__( - self, - reme: ReMe, - reme_model_name: str = "qwen-flash", - eval_model_name: str = "qwen3-max", - algo_version: str = "v1", - enable_thinking_params: bool = False, - ): - self.reme = reme - self.reme_model_name = reme_model_name - self.eval_model_name = eval_model_name - self.algo_version = algo_version - self.enable_thinking_params = enable_thinking_params - - async def search_memory( - self, - query: str, - user_id: str, - top_k: int = 20, - ) -> tuple[dict, list, float]: - """ - Search memory using ReMe and return structured answer with reasoning. - - Returns: - tuple: (answer_dict, agent_messages, duration_ms) - answer_dict contains: {"reasoning": str, "answer": str, "memories": str} - """ - start = time.time() - - # Retrieve memories from ReMe using new API - result = await self.reme.retrieve_memory( - llm_config_name="qwen3-max", - query=query, - retrieve_top_k=top_k, - user_name=user_id, - version=self.algo_version, - return_dict=True, - enable_time_filter=True, - enable_thinking_params=True, - ) - - # Extract memories from response - memories = result["answer"] - agent_messages = [x.simple_dump(enable_argument_dict=True) for x in result["messages"]] - retrieved_nodes = [x.model_dump(exclude_none=True) for x in result["retrieved_nodes"]] - - # Use LLM to generate structured answer from memories - answer_result = await answer_question_with_memories( - reme=self.reme, - question=query, - memories=memories, - user_id=user_id, - model_name=self.eval_model_name, - ) - - # Add original memories to the result - answer_result["memories"] = memories - answer_result["retrieved_nodes"] = retrieved_nodes - - duration_ms = (time.time() - start) * 1000 - return answer_result, agent_messages, duration_ms - - -# ==================== Answer Judge ==================== - - -class LongMemEvalJudge: - """LongMemEval answer judge using LLM.""" - - def __init__(self, reme: ReMe, model: str = "qwen3-max"): - self.reme = reme - self.model = model - - async def judge_answer( - self, - question_type: str, - question: str, - answer: str, - response: str, - abstention: bool = False, - ) -> dict: - """ - Judge if the model's response is correct. - - Returns: - dict with is_correct, llm_response, and judge_prompt - """ - prompt = get_anscheck_prompt(question_type, question, answer, response, abstention) - - try: - llm_response = await self.reme.get_llm("default").simple_request( - prompt=prompt, - model_name=self.model, - ) - llm_response_lower = llm_response.strip().lower() - is_correct = llm_response_lower.startswith("yes") - - return { - "is_correct": is_correct, - "llm_response": llm_response, - "judge_prompt": prompt, - } - except Exception as e: - return { - "is_correct": None, - "error": str(e), - "judge_prompt": prompt, - } - - -# ==================== Metrics ==================== - - -class MetricsAggregator: - """Aggregates evaluation metrics for LongMemEval.""" - - @staticmethod - def compute_metrics(results: list[dict]) -> dict[str, Any]: - """Compute overall and per-type metrics.""" - total = len(results) - correct = sum(1 for r in results if r.get("judgment", {}).get("is_correct") is True) - incorrect = sum(1 for r in results if r.get("judgment", {}).get("is_correct") is False) - error = total - correct - incorrect - - metrics = { - "total": total, - "correct": correct, - "incorrect": incorrect, - "error": error, - "accuracy": correct / total if total > 0 else 0, - "accuracy_valid": correct / (correct + incorrect) if (correct + incorrect) > 0 else 0, - } - - # Per question type statistics - type_stats = {} - for r in results: - qtype = r.get("question_type", "unknown") - if qtype not in type_stats: - type_stats[qtype] = {"total": 0, "correct": 0, "incorrect": 0} - type_stats[qtype]["total"] += 1 - if r.get("judgment", {}).get("is_correct") is True: - type_stats[qtype]["correct"] += 1 - elif r.get("judgment", {}).get("is_correct") is False: - type_stats[qtype]["incorrect"] += 1 - - metrics["by_question_type"] = { - qtype: { - **stats, - "accuracy": stats["correct"] / stats["total"] if stats["total"] > 0 else 0, - "accuracy_valid": ( - stats["correct"] / (stats["correct"] + stats["incorrect"]) - if (stats["correct"] + stats["incorrect"]) > 0 - else 0 - ), - } - for qtype, stats in type_stats.items() - } - - return metrics - - @staticmethod - def compute_timing_stats(results: list[dict]) -> dict[str, Any]: - """Compute timing statistics.""" - retrieve_times = [] - - for r in results: - retrieve_ms = r.get("retrieve_duration_ms", 0) - if retrieve_ms > 0: - retrieve_times.append(retrieve_ms) - - def compute_stats(times: list[float]) -> dict: - if not times: - return {"count": 0, "total_ms": 0, "avg_ms": 0, "min_ms": 0, "max_ms": 0} - - return { - "count": len(times), - "total_ms": sum(times), - "total_min": sum(times) / 1000 / 60, - "avg_ms": sum(times) / len(times), - "min_ms": min(times), - "max_ms": max(times), - } - - return { - "retrieve": compute_stats(retrieve_times), - "total_time_min": sum(retrieve_times) / 1000 / 60, - } - - -# ==================== Main Pipeline ==================== - - -class LongMemEvalRetrieveEvaluator: - """Retrieve-only evaluator for LongMemEval benchmark using ReMe.""" - - def __init__(self, config: RetrieveEvalConfig): - self.config = config - self.reme = ReMe( - default_llm_config={ - "model_name": self.config.reme_model_name, - }, - llms={ - "qwen-plus-think": { - "backend": "openai", - "model_name": "qwen-plus", - "extra_body": { - "enable_thinking": True, - }, - }, - "qwen3-max-think": { - "backend": "openai", - "model_name": "qwen3-max", - "extra_body": { - "enable_thinking": True, - }, - }, - "qwen3-max": { - "backend": "openai", - "model_name": "qwen3-max", - "extra_body": { - "enable_thinking": False, - }, - }, - }, - ) - - # Load evaluation prompts into ReMe's prompt handler - prompts_yaml_path = Path(__file__).parent / "eval_reme.yaml" - self.reme.prompt_handler.load_prompt_by_file(prompts_yaml_path) - - self.file_manager = FileManager(config.output_dir) - self.retrieve_processor = RetrieveProcessor( - self.reme, - config.reme_model_name, - config.eval_model_name, - config.algo_version, - config.enable_thinking_params, - ) - self.judge = LongMemEvalJudge(self.reme, config.eval_model_name) - self.data_loader = DataLoader() - - async def __aenter__(self): - """Async context manager entry.""" - await self.reme.start() - return self - - async def __aexit__(self, exc_type, exc_val, exc_tb): - """Async context manager exit with cleanup.""" - await self.reme.close() - return False - - async def process_question_entry(self, entry: dict, idx: int) -> dict: - """Process a single question entry (retrieve + judge only). - - Args: - entry: A question entry from LongMemEval dataset - idx: Index of the question - - Returns: - Result dictionary - """ - question_id = entry["question_id"] - question = entry["question"] - answer = entry["answer"] - question_type = entry["question_type"] - question_date = entry.get("question_date", "") - haystack_dates = entry["haystack_dates"] - haystack_session_ids = entry["haystack_session_ids"] - haystack_sessions = entry["haystack_sessions"] - - # Use question_id as user_id for isolation (same as full eval) - user_id = f"longmemeval_{question_id}" - - logger.info(f"\n{'='*60}") - logger.info(f"Question ID: {question_id}") - logger.info(f"Question Type: {question_type}") - logger.info(f"Question: {question}") - logger.info(f"Question_date: {question_date}") - logger.info(f"Answer: {answer}") - logger.info(f"Number of sessions: {len(haystack_sessions)}") - logger.info(f"{'='*60}") - - # Skip summary phase - directly search memory and answer question - logger.info(" Retrieving and answering question using ReMe...") - answer_dict, retrieve_messages, retrieve_duration_ms = await self.retrieve_processor.search_memory( - query=f"[Question_date: {question_date} | Question_type: {question_type}] " + question, - user_id=user_id, - top_k=self.config.top_k, - ) - - # Extract answer and reasoning from the structured response - model_response = answer_dict.get("answer", "") - model_reasoning = answer_dict.get("reasoning", "") - retrieved_memories = answer_dict.get("memories", "") - retrieved_nodes = answer_dict.get("retrieved_nodes", []) - - # Judge answer correctness - logger.info(" Judging answer correctness...") - judgment = await self.judge.judge_answer( - question_type=question_type, - question=question, - answer=answer, - response=model_response, - ) - - is_correct = judgment.get("is_correct") - logger.info( - f" → Answer judgment: {'Correct' if is_correct else 'Incorrect' if is_correct is False else 'Error'}", - ) - - result = { - "question_id": question_id, - "question_type": question_type, - "question": question, - "answer": answer, - "question_date": question_date, - "haystack_dates": haystack_dates, - "haystack_session_ids": haystack_session_ids, - "num_sessions": len(haystack_sessions), - "model_response": model_response, - "model_reasoning": model_reasoning, - "retrieved_memories": retrieved_memories, - "retrieved_nodes": retrieved_nodes, - "judgment": judgment, - "retrieve_duration_ms": retrieve_duration_ms, - "retrieve_messages": retrieve_messages, - } - - # Save individual result - self.file_manager.save_question_result(idx, question_id, result) - - logger.info(f" Question {question_id} - Completed") - - return result - - async def run_evaluation(self): - """Run the retrieve-only evaluation pipeline with parallel processing.""" - start_time = time.time() - - # NOTE: Do NOT clear vector store - we assume memories are already there from previous summary run - - # Load dataset - logger.info(f"Loading dataset from: {self.config.data_path}") - all_data = self.data_loader.load_json(self.config.data_path) - logger.info(f"Total questions in dataset: {len(all_data)}") - - # Filter by question type - logger.info(f"Filtering by type (samples_per_type={self.config.samples_per_type}):") - filtered_data = self.data_loader.filter_by_type(all_data, self.config.samples_per_type) - logger.info(f"Selected {len(filtered_data)} questions after filtering") - - # Apply start_index and end_index on filtered data - end_index = self.config.end_index or len(filtered_data) - start_index = self.config.start_index - end_index = min(end_index, len(filtered_data)) - - # Get the slice we want to process - data_to_process = filtered_data[start_index:end_index] - total_questions = len(data_to_process) - - logger.info(f"Processing {total_questions} questions (index {start_index} to {end_index - 1})") - - print("\n" + "=" * 80) - print("LONGMEMEVAL EVALUATION - REME (RETRIEVE ONLY)") - print(f"Samples per type: {self.config.samples_per_type} (-1 = all)") - print(f"Questions to process: {total_questions} | Top-K: {self.config.top_k}") - print(f"Max Concurrency: {self.config.max_concurrency}") - print(f"ReMe Model: {self.config.reme_model_name} | Eval Model: {self.config.eval_model_name}") - print(f"Algo Version: {self.config.algo_version}") - print("⚠️ NOTE: Assumes memories are already in vector store from previous summary run") - print("=" * 80 + "\n") - - # Use semaphore to control concurrency - semaphore = asyncio.Semaphore(self.config.max_concurrency) - - async def process_with_semaphore(idx: int, original_idx: int, entry: dict) -> Optional[dict]: - """Process a question with semaphore for concurrency control.""" - async with semaphore: - question_id = entry["question_id"] - - # Check cache first (use original index for cache file naming) - cached_result = self.file_manager.load_question_result(original_idx, question_id) - if cached_result: - print(f"⚡ [{idx}/{total_questions}] Skipping question {original_idx} (cached)") - return cached_result - - print(f"\n{'#'*60}") - print(f"### [{idx}/{total_questions}] Processing Question {original_idx} ###") - print(f"{'#'*60}") - - try: - result = await self.process_question_entry(entry, original_idx) - print(f"✅ [{idx}/{total_questions}] Completed question {original_idx}") - return result - except Exception as e: - logger.error(f"❌ Error processing question {original_idx}: {e}") - import traceback - - traceback.print_exc() - return { - "question_id": question_id, - "error": str(e), - "question_type": entry.get("question_type", "unknown"), - "question": entry.get("question", ""), - "answer": entry.get("answer", ""), - "judgment": {"is_correct": None, "error": str(e)}, - } - - # Create all tasks from filtered data (each item is a tuple of (original_idx, entry)) - tasks = [ - process_with_semaphore(idx + 1, original_idx, entry) - for idx, (original_idx, entry) in enumerate(data_to_process) - ] - - # Execute in parallel with controlled concurrency - all_results = await asyncio.gather(*tasks, return_exceptions=False) - - # Filter out None results if any - all_results = [r for r in all_results if r is not None] - - # Save summary - self.file_manager.save_summary(all_results) - - elapsed = time.time() - start_time - print(f"\n✅ Processing completed in {elapsed:.2f}s") - if total_questions > 0: - print(f" Average time per question: {elapsed / total_questions:.2f}s") - - # Compute and report metrics - self._report_metrics(all_results) - - return all_results - - def _report_metrics(self, results: list[dict]): - """Report evaluation metrics.""" - metrics = MetricsAggregator.compute_metrics(results) - timing_stats = MetricsAggregator.compute_timing_stats(results) - - print("\n" + "=" * 80) - print("EVALUATION SUMMARY - LONGMEMEVAL - REME (RETRIEVE ONLY)") - print("=" * 80 + "\n") - - print("📊 Overall Results:") - print(f" ✅ Correct: {metrics['correct']}/{metrics['total']} ({100*metrics['accuracy']:.2f}%)") - print( - f" ❌ Incorrect: {metrics['incorrect']}/{metrics['total']}" - f" ({100*metrics['incorrect']/metrics['total'] if metrics['total'] > 0 else 0:.2f}%)", - ) - if metrics["error"] > 0: - print(f" ⚠️ Error: {metrics['error']}/{metrics['total']} ({100*metrics['error']/metrics['total']:.2f}%)") - print(f" Accuracy (valid): {100*metrics['accuracy_valid']:.2f}%") - - print("\n📊 Accuracy by Question Type:") - print("-" * 60) - print(f"{'Question Type':<30} {'Correct':<10} {'Total':<10} {'Accuracy':<10}") - print("-" * 60) - for qtype in sorted(metrics["by_question_type"].keys()): - stats = metrics["by_question_type"][qtype] - print(f"{qtype:<30} {stats['correct']:<10} {stats['total']:<10} {100*stats['accuracy']:.2f}%") - print("-" * 60) - - print("\n⏱️ Timing Statistics (Retrieve Only):") - retrieve = timing_stats["retrieve"] - print(" Memory Retrieval:") - print(f" Total Time: {retrieve['total_min']:.2f} min") - print(f" Avg per Q: {retrieve['avg_ms']:.0f} ms") - print(f" Total Time: {timing_stats['total_time_min']:.2f} min") - - # Save metrics - final_results = { - "accuracy": metrics, - "timing": timing_stats, - } - metrics_file = self.file_manager.base_dir / "eval_statistics.json" - with open(metrics_file, "w", encoding="utf-8") as f: - json.dump(final_results, f, indent=4, ensure_ascii=False) - print(f"\n📁 Statistics saved to: {metrics_file}") - - print("\n" + "=" * 80) - - -# ==================== Entry Point ==================== - - -async def main_async( - data_path: str, - top_k: int = 20, - start_index: int = 0, - end_index: Optional[int] = None, - max_concurrency: int = 1, - output_dir: str = "bench_results/longmemeval_reme_retrieve", - reme_model_name: str = "qwen-flash", - eval_model_name: str = "qwen3-max", - algo_version: str = "v1", - samples_per_type: int = -1, - enable_thinking_params: bool = False, - summary_results_dir: Optional[str] = None, -): - """Main async entry point for LongMemEval retrieve-only evaluation with proper resource cleanup.""" - config = RetrieveEvalConfig( - data_path=data_path, - top_k=top_k, - start_index=start_index, - end_index=end_index, - max_concurrency=max_concurrency, - output_dir=output_dir, - reme_model_name=reme_model_name, - eval_model_name=eval_model_name, - algo_version=algo_version, - samples_per_type=samples_per_type, - enable_thinking_params=enable_thinking_params, - summary_results_dir=summary_results_dir, - ) - - # Use async context manager for automatic cleanup - async with LongMemEvalRetrieveEvaluator(config) as evaluator: - await evaluator.run_evaluation() - - -def main( - data_path: str, - top_k: int = 20, - start_index: int = 0, - end_index: Optional[int] = None, - max_concurrency: int = 1, - output_dir: str = "bench_results/longmemeval_reme_retrieve", - reme_model_name: str = "qwen-flash", - eval_model_name: str = "qwen3-max", - algo_version: str = "v1", - samples_per_type: int = -1, - enable_thinking_params: bool = False, - summary_results_dir: Optional[str] = None, -): - """Main entry point for LongMemEval retrieve-only evaluation.""" - asyncio.run( - main_async( - data_path=data_path, - top_k=top_k, - start_index=start_index, - end_index=end_index, - max_concurrency=max_concurrency, - output_dir=output_dir, - reme_model_name=reme_model_name, - eval_model_name=eval_model_name, - algo_version=algo_version, - samples_per_type=samples_per_type, - enable_thinking_params=enable_thinking_params, - summary_results_dir=summary_results_dir, - ), - ) - - -if __name__ == "__main__": - import argparse - - parser = argparse.ArgumentParser( - description="Evaluate ReMe on LongMemEval benchmark (Retrieve Phase Only)", - ) - parser.add_argument( - "--data_path", - type=str, - # default="/Users/zhouwk/PycharmProjects/MemAgent/dataset/longmemeval/longmemeval_s_cleaned.json", - default="/Users/zhouwk/PycharmProjects/MemAgent/dataset/longmemeval/longmemeval_oracle.json", - help="Path to LongMemEval JSON file", - ) - parser.add_argument( - "--top_k", - type=int, - default=10, - help="Number of memories to retrieve (default: 10)", - ) - parser.add_argument( - "--start_index", - type=int, - default=0, - help="Start index for processing questions (default: 0)", - ) - parser.add_argument( - "--end_index", - type=int, - default=None, - help="End index for processing questions (default: None, process all)", - ) - parser.add_argument( - "--max_concurrency", - type=int, - default=8, - help="Maximum concurrent question processing (default: 1)", - ) - parser.add_argument( - "--output_dir", - type=str, - default="bench_results/longmemeval_reme_retrieve_gpt4", - help="Output directory for results", - ) - parser.add_argument( - "--reme_model_name", - type=str, - default="gpt-4o-mini-2024-07-18", - help="Model name for ReMe operations (default: gpt-4o-mini-2024-07-18)", - ) - parser.add_argument( - "--eval_model_name", - type=str, - default="gpt-4o-mini-2024-07-18", - help="Model name for evaluation/judgment (default: gpt-4o-mini-2024-07-18)", - ) - parser.add_argument( - "--algo_version", - type=str, - default="longmemeval", - help="Algorithm version for retrieval (default: longmemeval)", - ) - parser.add_argument( - "--samples_per_type", - type=int, - default=4, - help="Number of samples per question type, -1 for all (default: 4)", - ) - parser.add_argument( - "--enable_thinking_params", - action="store_true", - default=False, - help="Enable thinking parameters for retrieval (default: False)", - ) - parser.add_argument( - "--summary_results_dir", - type=str, - default="/Users/zhouwk/PycharmProjects/ReMe/benchmark/longmemeval/bench_results", - help="Optional: path to previous summary results directory (for reference)", - ) - parser.add_argument( - "--no_cache", - action="store_true", - default=False, - help="Ignore cached results and re-run all questions (default: False)", - ) - - args = parser.parse_args() - print(f"args={args}!") - - main( - data_path=args.data_path, - top_k=args.top_k, - start_index=args.start_index, - end_index=args.end_index, - max_concurrency=args.max_concurrency, - output_dir=args.output_dir, - reme_model_name=args.reme_model_name, - eval_model_name=args.eval_model_name, - algo_version=args.algo_version, - samples_per_type=args.samples_per_type, - enable_thinking_params=args.enable_thinking_params, - summary_results_dir=args.summary_results_dir, - ) diff --git a/benchmark/longmemeval/eval_tools.py b/benchmark/longmemeval/eval_tools.py deleted file mode 100644 index a650d716..00000000 --- a/benchmark/longmemeval/eval_tools.py +++ /dev/null @@ -1,232 +0,0 @@ -"""Evaluation tools for ReMe LongMemEval benchmark.""" - -from pathlib import Path - -import yaml - -from reme.reme import ReMe - - -# Load prompts from YAML file -_YAML_PATH = Path(__file__).parent / "eval_reme.yaml" -with open(_YAML_PATH, "r", encoding="utf-8") as f: - _PROMPTS = yaml.safe_load(f) - - -async def evaluation_for_memory_integrity( - reme: ReMe, - extract_memories: str, - target_memory: str, - model_name: str = "qwen3-max", -) -> dict: - """ - Memory Integrity Evaluation - - Args: - reme: ReMe instance - extract_memories: A formatted string concatenating all memory points extracted by the memory system. - target_memory: The target key memory point. - model_name: Model name for evaluation - - Returns: - dict with 'reasoning' and 'score' fields - """ - prompt = _PROMPTS["EVALUATION_PROMPT_FOR_MEMORY_INTEGRITY"].format( - memories=extract_memories, - expected_memory_point=target_memory, - ) - - result = await reme.llm.simple_request_for_json( - prompt=prompt, - model_name=model_name, - ) - - return result - - -async def evaluation_for_memory_accuracy( - reme: ReMe, - dialogue: str, - golden_memories: str, - candidate_memory: str, - model_name: str = "qwen3-max", -) -> dict: - """ - Memory Accuracy Evaluation - - Args: - reme: ReMe instance - dialogue: The complete human-machine dialogue record. - golden_memories: The core memory points for this dialogue segment in the evaluation set . - candidate_memory: A specific memory point extracted by the memory system being evaluated. - model_name: Model name for evaluation - - Returns: - dict with 'accuracy_score', 'is_included_in_golden_memories', and 'reason' fields - """ - prompt = _PROMPTS["EVALUATION_PROMPT_FOR_MEMORY_ACCURACY"].format( - dialogue=dialogue, - golden_memories=golden_memories, - candidate_memory=candidate_memory, - ) - - result = await reme.llm.simple_request_for_json( - prompt=prompt, - model_name=model_name, - ) - - return result - - -async def evaluation_for_update_memory( - reme: ReMe, - extract_memories: str, - target_update_memory: str, - original_memory: str, - model_name: str = "qwen3-max", -) -> dict: - """ - Memory Update Evaluation - - Args: - reme: ReMe instance - extract_memories: A formatted string concatenating all memory points extracted by the memory system . - target_update_memory: The target updated memory point. - original_memory: A formatted string concatenating all original memory points corresponding. - model_name: Model name for evaluation - - Returns: - dict with 'reason' and 'evaluation_result' fields - """ - prompt = _PROMPTS["EVALUATION_PROMPT_FOR_UPDATE_MEMORY"].format( - memories=extract_memories, - updated_memory=target_update_memory, - original_memory=original_memory, - ) - - result = await reme.llm.simple_request_for_json( - prompt=prompt, - model_name=model_name, - ) - - return result - - -async def evaluation_for_question( - reme: ReMe, - question: str, - reference_answer: str, - key_memory_points: str, - response: str, - model_name: str = "qwen3-max", -) -> dict: - """ - Question-Answering Evaluation - - Args: - reme: ReMe instance - question: The question string to be evaluated. - reference_answer: The reference (gold-standard) answer. - key_memory_points: The memory points used to derive the reference answer. - response: The answer produced by the memory system. - model_name: Model name for evaluation - - Returns: - dict with 'reasoning' and 'evaluation_result' fields - """ - prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION"].format( - question=question, - reference_answer=reference_answer, - key_memory_points=key_memory_points, - response=response, - ) - - result = await reme.llm.simple_request_for_json( - prompt=prompt, - model_name=model_name, - ) - - return result - - -async def evaluation_for_question2( - reme: ReMe, - question: str, - reference_answer: str, - key_memory_points: str, - response: str, - dialogue: str = "", - model_name: str = "qwen3-max", -) -> dict: - """ - Question-Answering Evaluation with Dialogue Context (Version 2) - - Args: - reme: ReMe instance - question: The question string to be evaluated. - reference_answer: The reference (gold-standard) answer. - key_memory_points: The memory points used to derive the reference answer. - response: The answer produced by the memory system. - dialogue: The formatted dialogue history (role, content, time_created). - model_name: Model name for evaluation - - Returns: - dict with 'reasoning' and 'evaluation_result' fields - """ - prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"].format( - question=question, - reference_answer=reference_answer, - key_memory_points=key_memory_points, - response=response, - dialogue=dialogue if dialogue else "", - ) - - result = await reme.llm.simple_request_for_json( - prompt=prompt, - model_name=model_name, - ) - - return result - - -async def answer_question_with_memories( - reme: ReMe, - question: str, - memories: str, - user_id: str = None, - model_name: str = "qwen3-max", -) -> dict: - """ - Answer a question using retrieved memories with PROMPT_MEMZERO_JSON template. - - Args: - reme: ReMe instance - question: The question to answer - memories: The retrieved memories (formatted as context) - user_id: Optional user ID for context formatting - model_name: Model name for LLM request - - Returns: - dict with 'reasoning' and 'answer' fields - """ - # Format context with memories - if user_id: - context = _PROMPTS["TEMPLATE_MEMOS"].format( - user_id=user_id, - memories=memories, - ) - else: - context = f"Memories:\n{memories}" - - # Use PROMPT_MEMZERO_JSON template for structured JSON response - prompt = _PROMPTS["PROMPT_MEMZERO_JSON"].format( - context=context, - question=question, - ) - - result = await reme.llm.simple_request_for_json( - prompt=prompt, - model_name=model_name, - ) - - return result diff --git a/benchmark/longmemeval/llms.py b/benchmark/longmemeval/llms.py deleted file mode 100644 index 8498b52f..00000000 --- a/benchmark/longmemeval/llms.py +++ /dev/null @@ -1,97 +0,0 @@ -"""LLM utilities for LongMemEval benchmark evaluation.""" - -import asyncio -import json -import logging -import re - -from tenacity import retry, stop_after_attempt, wait_random_exponential, before_sleep_log - -from reme.core.schema import Message -from reme.core.utils import load_env -from reme.reme import ReMe - -logger = logging.getLogger(__name__) - -load_env() - -WAIT_TIME_LOWER = 1 -WAIT_TIME_UPPER = 60 -RETRY_TIMES = 5 - - -@retry( - wait=wait_random_exponential(min=WAIT_TIME_LOWER, max=WAIT_TIME_UPPER), - stop=stop_after_attempt(3), - reraise=True, - before_sleep=before_sleep_log(logger, logging.WARNING), -) -async def llm_request(reme: ReMe, prompt: str, model_name: str = "qwen3-max", **kwargs) -> str: - """Make an LLM request using ReMe's LLM with optional model override. - - Args: - reme: ReMe instance - prompt: The prompt to send to the LLM - model_name: Optional model name to override the default model (default: "qwen3-max") - **kwargs: Additional arguments to pass to the chat method - - Returns: - The assistant's response content - """ - assistant_message = await reme.llm.chat( - messages=[ - Message(role="user", content=prompt), - ], - model_name=model_name, - **kwargs, - ) - return assistant_message.content - - -@retry( - wait=wait_random_exponential(min=WAIT_TIME_LOWER, max=WAIT_TIME_UPPER), - stop=stop_after_attempt(RETRY_TIMES), - reraise=True, - before_sleep=before_sleep_log(logger, logging.WARNING), -) -async def llm_request_for_json(reme: ReMe, prompt: str, model_name: str = "qwen-flash", **kwargs) -> dict: - """Make an LLM request expecting JSON response using ReMe's LLM. - - Args: - reme: ReMe instance - prompt: The prompt to send to the LLM - model_name: Optional model name to override the default model (default: "qwen-flash") - **kwargs: Additional arguments to pass to the chat method - - Returns: - Parsed JSON object from the LLM response - - Raises: - ValueError: If no JSON block is found in the model output - """ - content = await llm_request(reme, prompt, model_name=model_name, **kwargs) - - match = re.search(r"```json\s*(\{.*?\})\s*```", content, re.DOTALL) - if not match: - raise ValueError(f"No JSON block found in model output: {content}") - - json_str = match.group(1).strip() - return json.loads(json_str) - - -if __name__ == "__main__": - - async def test(): - """Simple manual test for JSON LLM request.""" - reme = ReMe() - await reme.start() - try: - r = await llm_request_for_json( - reme, - 'hello? answer in ```json\n{"answer": "..."}```', - ) - print(r) - finally: - await reme.close() - - asyncio.run(test()) diff --git a/reme/__init__.py b/reme/__init__.py index 604c0014..e4db3de0 100644 --- a/reme/__init__.py +++ b/reme/__init__.py @@ -1,23 +1,26 @@ -"""ReMe""" +"""ReMe CLI package.""" + +__version__ = "0.4.0.0" from . import config -from . import core -from . import extension -from . import memory +from . import constants +from . import enumeration +from . import schema +from . import steps +from . import utils +from .application import Application +from .components import BaseComponent from .reme import ReMe -__version__ = "0.3.1.10" - __all__ = [ - "config", - "core", - "extension", - "memory", + "Application", + "BaseComponent", "ReMe", + # submodules + "config", + "constants", + "enumeration", + "schema", + "steps", + "utils", ] - -""" -conda create -n fl_test2 python=3.10 -conda activate fl_test2 -conda env remove -n fl_test2 -""" diff --git a/reme4/application.py b/reme/application.py similarity index 100% rename from reme4/application.py rename to reme/application.py diff --git a/reme4/components/__init__.py b/reme/components/__init__.py similarity index 100% rename from reme4/components/__init__.py rename to reme/components/__init__.py diff --git a/reme4/components/agent_wrapper/__init__.py b/reme/components/agent_wrapper/__init__.py similarity index 100% rename from reme4/components/agent_wrapper/__init__.py rename to reme/components/agent_wrapper/__init__.py diff --git a/reme4/components/agent_wrapper/as_agent_wrapper.py b/reme/components/agent_wrapper/as_agent_wrapper.py similarity index 100% rename from reme4/components/agent_wrapper/as_agent_wrapper.py rename to reme/components/agent_wrapper/as_agent_wrapper.py diff --git a/reme4/components/agent_wrapper/base_agent_wrapper.py b/reme/components/agent_wrapper/base_agent_wrapper.py similarity index 100% rename from reme4/components/agent_wrapper/base_agent_wrapper.py rename to reme/components/agent_wrapper/base_agent_wrapper.py diff --git a/reme4/components/agent_wrapper/cc_agent_wrapper.py b/reme/components/agent_wrapper/cc_agent_wrapper.py similarity index 100% rename from reme4/components/agent_wrapper/cc_agent_wrapper.py rename to reme/components/agent_wrapper/cc_agent_wrapper.py diff --git a/reme4/components/application_context.py b/reme/components/application_context.py similarity index 100% rename from reme4/components/application_context.py rename to reme/components/application_context.py diff --git a/reme4/components/as_embedding/__init__.py b/reme/components/as_embedding/__init__.py similarity index 100% rename from reme4/components/as_embedding/__init__.py rename to reme/components/as_embedding/__init__.py diff --git a/reme4/components/as_llm/__init__.py b/reme/components/as_llm/__init__.py similarity index 100% rename from reme4/components/as_llm/__init__.py rename to reme/components/as_llm/__init__.py diff --git a/reme4/components/base_component.py b/reme/components/base_component.py similarity index 100% rename from reme4/components/base_component.py rename to reme/components/base_component.py diff --git a/reme4/components/client/__init__.py b/reme/components/client/__init__.py similarity index 100% rename from reme4/components/client/__init__.py rename to reme/components/client/__init__.py diff --git a/reme4/components/client/base_client.py b/reme/components/client/base_client.py similarity index 100% rename from reme4/components/client/base_client.py rename to reme/components/client/base_client.py diff --git a/reme4/components/client/http_client.py b/reme/components/client/http_client.py similarity index 100% rename from reme4/components/client/http_client.py rename to reme/components/client/http_client.py diff --git a/reme4/components/client/mcp_client.py b/reme/components/client/mcp_client.py similarity index 100% rename from reme4/components/client/mcp_client.py rename to reme/components/client/mcp_client.py diff --git a/reme4/components/component_registry.py b/reme/components/component_registry.py similarity index 100% rename from reme4/components/component_registry.py rename to reme/components/component_registry.py diff --git a/reme4/components/embedding_store/__init__.py b/reme/components/embedding_store/__init__.py similarity index 100% rename from reme4/components/embedding_store/__init__.py rename to reme/components/embedding_store/__init__.py diff --git a/reme4/components/embedding_store/base_embedding_store.py b/reme/components/embedding_store/base_embedding_store.py similarity index 100% rename from reme4/components/embedding_store/base_embedding_store.py rename to reme/components/embedding_store/base_embedding_store.py diff --git a/reme4/components/embedding_store/local_embedding_store.py b/reme/components/embedding_store/local_embedding_store.py similarity index 100% rename from reme4/components/embedding_store/local_embedding_store.py rename to reme/components/embedding_store/local_embedding_store.py diff --git a/reme4/components/file_catalog/__init__.py b/reme/components/file_catalog/__init__.py similarity index 100% rename from reme4/components/file_catalog/__init__.py rename to reme/components/file_catalog/__init__.py diff --git a/reme4/components/file_catalog/base_file_catalog.py b/reme/components/file_catalog/base_file_catalog.py similarity index 100% rename from reme4/components/file_catalog/base_file_catalog.py rename to reme/components/file_catalog/base_file_catalog.py diff --git a/reme4/components/file_catalog/local_file_catalog.py b/reme/components/file_catalog/local_file_catalog.py similarity index 100% rename from reme4/components/file_catalog/local_file_catalog.py rename to reme/components/file_catalog/local_file_catalog.py diff --git a/reme4/components/file_chunker/__init__.py b/reme/components/file_chunker/__init__.py similarity index 100% rename from reme4/components/file_chunker/__init__.py rename to reme/components/file_chunker/__init__.py diff --git a/reme4/components/file_chunker/base_file_chunker.py b/reme/components/file_chunker/base_file_chunker.py similarity index 100% rename from reme4/components/file_chunker/base_file_chunker.py rename to reme/components/file_chunker/base_file_chunker.py diff --git a/reme4/components/file_chunker/default_file_chunker.py b/reme/components/file_chunker/default_file_chunker.py similarity index 100% rename from reme4/components/file_chunker/default_file_chunker.py rename to reme/components/file_chunker/default_file_chunker.py diff --git a/reme4/components/file_chunker/markdown_file_chunker.py b/reme/components/file_chunker/markdown_file_chunker.py similarity index 100% rename from reme4/components/file_chunker/markdown_file_chunker.py rename to reme/components/file_chunker/markdown_file_chunker.py diff --git a/reme4/components/file_graph/__init__.py b/reme/components/file_graph/__init__.py similarity index 100% rename from reme4/components/file_graph/__init__.py rename to reme/components/file_graph/__init__.py diff --git a/reme4/components/file_graph/base_file_graph.py b/reme/components/file_graph/base_file_graph.py similarity index 100% rename from reme4/components/file_graph/base_file_graph.py rename to reme/components/file_graph/base_file_graph.py diff --git a/reme4/components/file_graph/local_file_graph.py b/reme/components/file_graph/local_file_graph.py similarity index 100% rename from reme4/components/file_graph/local_file_graph.py rename to reme/components/file_graph/local_file_graph.py diff --git a/reme4/components/file_graph/neo4j_file_graph.py b/reme/components/file_graph/neo4j_file_graph.py similarity index 100% rename from reme4/components/file_graph/neo4j_file_graph.py rename to reme/components/file_graph/neo4j_file_graph.py diff --git a/reme4/components/file_graph/nx_file_graph.py b/reme/components/file_graph/nx_file_graph.py similarity index 100% rename from reme4/components/file_graph/nx_file_graph.py rename to reme/components/file_graph/nx_file_graph.py diff --git a/reme4/components/file_store/__init__.py b/reme/components/file_store/__init__.py similarity index 100% rename from reme4/components/file_store/__init__.py rename to reme/components/file_store/__init__.py diff --git a/reme4/components/file_store/base_file_store.py b/reme/components/file_store/base_file_store.py similarity index 100% rename from reme4/components/file_store/base_file_store.py rename to reme/components/file_store/base_file_store.py diff --git a/reme4/components/file_store/faiss_local_file_store.py b/reme/components/file_store/faiss_local_file_store.py similarity index 100% rename from reme4/components/file_store/faiss_local_file_store.py rename to reme/components/file_store/faiss_local_file_store.py diff --git a/reme4/components/file_store/local_file_store.py b/reme/components/file_store/local_file_store.py similarity index 100% rename from reme4/components/file_store/local_file_store.py rename to reme/components/file_store/local_file_store.py diff --git a/reme4/components/job/__init__.py b/reme/components/job/__init__.py similarity index 100% rename from reme4/components/job/__init__.py rename to reme/components/job/__init__.py diff --git a/reme4/components/job/background_job.py b/reme/components/job/background_job.py similarity index 100% rename from reme4/components/job/background_job.py rename to reme/components/job/background_job.py diff --git a/reme4/components/job/base_job.py b/reme/components/job/base_job.py similarity index 100% rename from reme4/components/job/base_job.py rename to reme/components/job/base_job.py diff --git a/reme4/components/job/cron_job.py b/reme/components/job/cron_job.py similarity index 100% rename from reme4/components/job/cron_job.py rename to reme/components/job/cron_job.py diff --git a/reme4/components/job/stream_job.py b/reme/components/job/stream_job.py similarity index 100% rename from reme4/components/job/stream_job.py rename to reme/components/job/stream_job.py diff --git a/reme4/components/keyword_index/__init__.py b/reme/components/keyword_index/__init__.py similarity index 100% rename from reme4/components/keyword_index/__init__.py rename to reme/components/keyword_index/__init__.py diff --git a/reme4/components/keyword_index/base_keyword_index.py b/reme/components/keyword_index/base_keyword_index.py similarity index 100% rename from reme4/components/keyword_index/base_keyword_index.py rename to reme/components/keyword_index/base_keyword_index.py diff --git a/reme4/components/keyword_index/bm25_index.py b/reme/components/keyword_index/bm25_index.py similarity index 100% rename from reme4/components/keyword_index/bm25_index.py rename to reme/components/keyword_index/bm25_index.py diff --git a/reme4/components/prompt_handler.py b/reme/components/prompt_handler.py similarity index 100% rename from reme4/components/prompt_handler.py rename to reme/components/prompt_handler.py diff --git a/reme4/components/runtime_context.py b/reme/components/runtime_context.py similarity index 100% rename from reme4/components/runtime_context.py rename to reme/components/runtime_context.py diff --git a/reme4/components/service/__init__.py b/reme/components/service/__init__.py similarity index 100% rename from reme4/components/service/__init__.py rename to reme/components/service/__init__.py diff --git a/reme4/components/service/base_service.py b/reme/components/service/base_service.py similarity index 100% rename from reme4/components/service/base_service.py rename to reme/components/service/base_service.py diff --git a/reme4/components/service/http_service.py b/reme/components/service/http_service.py similarity index 100% rename from reme4/components/service/http_service.py rename to reme/components/service/http_service.py diff --git a/reme4/components/service/mcp_service.py b/reme/components/service/mcp_service.py similarity index 100% rename from reme4/components/service/mcp_service.py rename to reme/components/service/mcp_service.py diff --git a/reme4/components/tokenizer/__init__.py b/reme/components/tokenizer/__init__.py similarity index 100% rename from reme4/components/tokenizer/__init__.py rename to reme/components/tokenizer/__init__.py diff --git a/reme4/components/tokenizer/base_tokenizer.py b/reme/components/tokenizer/base_tokenizer.py similarity index 100% rename from reme4/components/tokenizer/base_tokenizer.py rename to reme/components/tokenizer/base_tokenizer.py diff --git a/reme4/components/tokenizer/jieba_tokenizer.py b/reme/components/tokenizer/jieba_tokenizer.py similarity index 100% rename from reme4/components/tokenizer/jieba_tokenizer.py rename to reme/components/tokenizer/jieba_tokenizer.py diff --git a/reme4/components/tokenizer/regex_tokenizer.py b/reme/components/tokenizer/regex_tokenizer.py similarity index 100% rename from reme4/components/tokenizer/regex_tokenizer.py rename to reme/components/tokenizer/regex_tokenizer.py diff --git a/reme/config/__init__.py b/reme/config/__init__.py index a8d954d2..c2903189 100644 --- a/reme/config/__init__.py +++ b/reme/config/__init__.py @@ -1,5 +1,8 @@ -"""config""" +"""Config""" -from .reme_config_parser import ReMeConfigParser +from .config_parser import parse_args, resolve_app_config -__all__ = ["ReMeConfigParser"] +__all__ = [ + "parse_args", + "resolve_app_config", +] diff --git a/reme4/config/config_parser.py b/reme/config/config_parser.py similarity index 100% rename from reme4/config/config_parser.py rename to reme/config/config_parser.py diff --git a/reme/config/reme_config_parser.py b/reme/config/reme_config_parser.py deleted file mode 100644 index 7a21f806..00000000 --- a/reme/config/reme_config_parser.py +++ /dev/null @@ -1,7 +0,0 @@ -"""Configuration parser for ReMe framework.""" - -from ..core.utils import PydanticConfigParser - - -class ReMeConfigParser(PydanticConfigParser): - """Configuration parser for ReMe framework.""" diff --git a/reme4/constants.py b/reme/constants.py similarity index 100% rename from reme4/constants.py rename to reme/constants.py diff --git a/reme/core/__init__.py b/reme/core/__init__.py deleted file mode 100644 index 726a3455..00000000 --- a/reme/core/__init__.py +++ /dev/null @@ -1,51 +0,0 @@ -"""Core""" - -from . import as_llm -from . import as_llm_formatter -from . import as_token_counter -from . import embedding -from . import enumeration -from . import file_store -from . import file_watcher -from . import flow -from . import llm -from . import op -from . import schema -from . import service -from . import token_counter -from . import utils -from . import vector_store -from .application import Application -from .base_dict import BaseDict -from .prompt_handler import PromptHandler -from .registry_factory import R, Registry, RegistryFactory -from .runtime_context import RuntimeContext -from .service_context import ServiceContext - -__all__ = [ - # Submodules - "as_llm", - "as_llm_formatter", - "as_token_counter", - "embedding", - "enumeration", - "file_watcher", - "flow", - "llm", - "file_store", - "op", - "schema", - "service", - "token_counter", - "utils", - "vector_store", - # Classes - "Application", - "BaseDict", - "PromptHandler", - "R", - "Registry", - "RegistryFactory", - "RuntimeContext", - "ServiceContext", -] diff --git a/reme/core/application.py b/reme/core/application.py deleted file mode 100644 index f776ed98..00000000 --- a/reme/core/application.py +++ /dev/null @@ -1,657 +0,0 @@ -"""High-level entry point for configuring and running ReMe services and flows.""" - -import asyncio -import os -from concurrent.futures import ThreadPoolExecutor -from pathlib import Path - -from .embedding import BaseEmbeddingModel -from .file_store import BaseFileStore -from .file_watcher import BaseFileWatcher -from .flow import BaseFlow -from .llm import BaseLLM -from .prompt_handler import PromptHandler -from .registry_factory import R -from .schema import ( - EmbeddingModelConfig, - Response, - ServiceConfig, - LLMConfig, - VectorStoreConfig, - FileStoreConfig, - FileWatcherConfig, - TokenCounterConfig, -) -from .service_context import ServiceContext -from .token_counter import BaseTokenCounter -from .utils import execute_stream_task, PydanticConfigParser, init_logger, MCPClient, print_logo, get_logger, load_env -from .vector_store import BaseVectorStore - -logger = get_logger() - - -class Application: - """Application wrapper that wires together service context, flows, and runtimes.""" - - def __init__( - self, - *args, - llm_api_key: str | None = None, - llm_base_url: str | None = None, - embedding_api_key: str | None = None, - embedding_base_url: str | None = None, - working_dir: str | None = None, - config_path: str | None = None, - enable_logo: bool = True, - log_to_console: bool = True, - log_to_file: bool = True, - enable_load_env: bool = True, - parser: type[PydanticConfigParser] | None = None, - default_as_llm_config: dict | None = None, - default_as_llm_formatter_config: dict | None = None, - default_llm_config: dict | None = None, - default_embedding_model_config: dict | None = None, - default_vector_store_config: dict | None = None, - default_file_store_config: dict | None = None, - default_token_counter_config: dict | None = None, - default_file_watcher_config: dict | None = None, - **kwargs, - ): - - if enable_load_env: - load_env() - - self.llm_api_key = llm_api_key or os.getenv("LLM_API_KEY", "") - self.llm_base_url = llm_base_url or os.getenv("LLM_BASE_URL", "") - self.embedding_api_key = embedding_api_key or os.getenv("EMBEDDING_API_KEY", "") - self.embedding_base_url = embedding_base_url or os.getenv("EMBEDDING_BASE_URL", "") - - self.service_context = ServiceContext( - *args, - service_config=None, - parser=parser, - working_dir=working_dir, - config_path=config_path, - enable_logo=enable_logo, - log_to_console=log_to_console, - log_to_file=log_to_file, - default_as_llm_config=default_as_llm_config, - default_as_llm_formatter_config=default_as_llm_formatter_config, - default_llm_config=default_llm_config, - default_embedding_model_config=default_embedding_model_config, - default_vector_store_config=default_vector_store_config, - default_file_store_config=default_file_store_config, - default_token_counter_config=default_token_counter_config, - default_file_watcher_config=default_file_watcher_config, - **kwargs, - ) - - self.prompt_handler = PromptHandler(language=self.service_config.language) - - # NOTE: flows are initialized here to start service! - self.init_flows() - - self._started: bool = False - - @classmethod - async def create(cls, *args, **kwargs) -> "Application": - """Create and start an Application instance asynchronously.""" - instance = cls(*args, **kwargs) - await instance.start() - return instance - - def init_flows(self): - """Initialize flows.""" - expression_flow_cls = None - for name, flow_cls in R.flows.items(): - if not self._filter_flows(name): - continue - - if name == "ExpressionFlow": - expression_flow_cls = flow_cls - else: - flow: "BaseFlow" = flow_cls(name=name, service_context=self.service_context) - self.service_context.flows[flow.name] = flow - - if expression_flow_cls is not None: - for name, flow_config in self.service_config.flows.items(): - if not self._filter_flows(name): - continue - flow_config.name = name - flow: BaseFlow = expression_flow_cls( # noqa - flow_config=flow_config, - service_context=self.service_context, - ) - self.service_context.flows[flow.name] = flow - else: - logger.info("No expression flow found, please check your configuration.") - - def _filter_flows(self, name: str) -> bool: - """Filter flows based on enabled_flows and disabled_flows configuration.""" - if self.service_config.enabled_flows: - return name in self.service_config.enabled_flows - elif self.service_config.disabled_flows: - return name not in self.service_config.disabled_flows - else: - return True - - @property - def service_config(self) -> ServiceConfig: - """Get the service configuration.""" - return self.service_context.service_config - - async def start(self): - """Start the service context by initializing all configured components.""" - if self._started: - logger.warning("Application has already started.") - return self - - init_logger( - log_to_console=self.service_config.log_to_console, - log_to_file=self.service_config.log_to_file, - ) - logger.info(f"Init ReMe with config: {self.service_config.model_dump_json()}") - - working_path = Path(self.service_config.working_dir) - working_path.mkdir(parents=True, exist_ok=True) - - if self.service_config.ray_max_workers > 1: - import ray - - if not ray.is_initialized(): - ray.init(num_cpus=self.service_config.ray_max_workers) - - if self.service_config.thread_pool_max_workers > 0 and ( - self.service_context.thread_pool is None - or self.service_context.thread_pool._shutdown # pylint: disable=protected-access - ): - self.service_context.thread_pool = ThreadPoolExecutor( - max_workers=self.service_config.thread_pool_max_workers, - ) - elif self.service_config.thread_pool_max_workers <= 0: - logger.info("Thread pool is disabled (thread_pool_max_workers <= 0)") - - if self.service_context.service_config.enable_logo: - print_logo(service_config=self.service_config) - - for name, config in self.service_config.as_llms.items(): - if config.backend not in R.as_llms: - logger.warning(f"AS LLM backend {config.backend} is not supported.") - else: - try: - config_dict = config.model_dump(exclude={"backend"}) - if not config_dict.get("api_key", ""): - config_dict["api_key"] = self.llm_api_key - if "client_kwargs" not in config_dict: - config_dict["client_kwargs"] = {} - if not config_dict["client_kwargs"].get("base_url", ""): - config_dict["client_kwargs"]["base_url"] = self.llm_base_url - self.service_context.as_llms[name] = R.as_llms[config.backend](**config_dict) - except Exception as e: - logger.error(f"Failed to initialize AS LLM '{name}': {e}") - - for name, config in self.service_config.as_llm_formatters.items(): - if config.backend not in R.as_llm_formatters: - logger.warning(f"AS LLM formatter backend {config.backend} is not supported.") - else: - try: - config_dict = config.model_dump(exclude={"backend"}) - self.service_context.as_llm_formatters[name] = R.as_llm_formatters[config.backend](**config_dict) - except Exception as e: - logger.error(f"Failed to initialize AS LLM formatter '{name}': {e}") - - for name, config in self.service_config.as_token_counters.items(): - if config.backend not in R.as_token_counters: - logger.warning(f"Token counter backend {config.backend} is not supported.") - else: - try: - config_dict = config.model_dump(exclude={"backend"}) - self.service_context.as_token_counters[name] = R.as_token_counters[config.backend](**config_dict) - except Exception as e: - logger.error(f"Failed to initialize AS token counter '{name}': {e}") - - for name, config in self.service_config.llms.items(): - if config.backend not in R.llms: - logger.warning(f"LLM backend {config.backend} is not supported.") - else: - config_dict = config.model_dump(exclude={"backend"}) - config_dict.setdefault("api_key", self.llm_api_key) - config_dict.setdefault("base_url", self.llm_base_url) - self.service_context.llms[name] = R.llms[config.backend](**config_dict) - await self.service_context.llms[name].start() - - for name, config in self.service_config.embedding_models.items(): - if config.backend not in R.embedding_models: - logger.warning(f"Embedding model backend {config.backend} is not supported.") - else: - config_dict = config.model_dump(exclude={"backend"}) - config_dict.setdefault("api_key", self.embedding_api_key) - config_dict.setdefault("base_url", self.embedding_base_url) - config_dict.setdefault("cache_dir", working_path / "embedding_cache") - self.service_context.embedding_models[name] = R.embedding_models[config.backend](**config_dict) - await self.service_context.embedding_models[name].start() - - for name, config in self.service_config.token_counters.items(): - if config.backend not in R.token_counters: - logger.warning(f"Token counter backend {config.backend} is not supported.") - else: - config_dict = config.model_dump(exclude={"backend"}) - self.service_context.token_counters[name] = R.token_counters[config.backend](**config_dict) - - for name, config in self.service_config.vector_stores.items(): - if config.backend not in R.vector_stores: - logger.warning(f"Vector store backend {config.backend} is not supported.") - else: - config_dict = config.model_dump(exclude={"backend", "embedding_model"}) - config_dict.update( - { - "embedding_model": self.service_context.embedding_models[config.embedding_model], - "db_path": working_path / "vector_store", - }, - ) - self.service_context.vector_stores[name] = R.vector_stores[config.backend](**config_dict) - await self.service_context.vector_stores[name].start() - - for name, config in self.service_config.file_stores.items(): - if config.backend not in R.file_stores: - logger.warning(f"File store backend {config.backend} is not supported.") - else: - config_dict = config.model_dump(exclude={"backend", "embedding_model"}) - config_dict.update( - { - "embedding_model": self.service_context.embedding_models[config.embedding_model], - "db_path": working_path / "file_store", - }, - ) - self.service_context.file_stores[name] = R.file_stores[config.backend](**config_dict) - await self.service_context.file_stores[name].start() - - for name, config in self.service_config.file_watchers.items(): - if config.backend not in R.file_watchers: - logger.warning(f"File watcher backend {config.backend} is not supported.") - else: - config_dict = config.model_dump(exclude={"backend", "file_store"}) - config_dict["file_store"] = self.service_context.file_stores[config.file_store] - self.service_context.file_watchers[name] = R.file_watchers[config.backend](**config_dict) - await self.service_context.file_watchers[name].start() - - if self.service_config.mcp_servers: - await self.prepare_mcp_servers() - - self._started = True - logger.info("ReMe Application started") - return self - - # pylint: disable=too-many-statements - async def restart(self, restart_config: dict): - """Restart the application with new config.""" - - working_path = Path(self.service_config.working_dir) - working_path.mkdir(parents=True, exist_ok=True) - - # as_llms - if "as_llms" in restart_config: - as_llms_config = restart_config["as_llms"] - assert isinstance(as_llms_config, dict) - for name, config in as_llms_config.items(): - if name in self.service_context.as_llms: - del self.service_context.as_llms[name] - - if config.get("backend") not in R.as_llms: - logger.warning(f"AS LLM backend {config.get('backend')} is not supported.") - continue - - try: - config_dict = {k: v for k, v in config.items() if k != "backend"} - if not config_dict.get("api_key", ""): - config_dict["api_key"] = self.llm_api_key - if "client_kwargs" not in config_dict: - config_dict["client_kwargs"] = {} - if not config_dict["client_kwargs"].get("base_url", ""): - config_dict["client_kwargs"]["base_url"] = self.llm_base_url - self.service_context.as_llms[name] = R.as_llms[config["backend"]](**config_dict) - logger.info(f"Restarted AS LLM: {name}") - except Exception as e: - logger.error(f"Failed to restart AS LLM '{name}': {e}") - - # as_llm_formatters - if "as_llm_formatters" in restart_config: - as_llm_formatters_config = restart_config["as_llm_formatters"] - assert isinstance(as_llm_formatters_config, dict) - for name, config in as_llm_formatters_config.items(): - if name in self.service_context.as_llm_formatters: - del self.service_context.as_llm_formatters[name] - - if config.get("backend") not in R.as_llm_formatters: - logger.warning(f"AS LLM formatter backend {config.get('backend')} is not supported.") - continue - try: - config_dict = {k: v for k, v in config.items() if k != "backend"} - self.service_context.as_llm_formatters[name] = R.as_llm_formatters[config["backend"]](**config_dict) - logger.info(f"Restarted AS LLM formatter: {name}") - except Exception as e: - logger.error(f"Failed to restart AS LLM formatter '{name}': {e}") - - # as_token_counters - if "as_token_counters" in restart_config: - as_token_counters_config = restart_config["as_token_counters"] - assert isinstance(as_token_counters_config, dict) - for name, config in as_token_counters_config.items(): - if name in self.service_context.as_token_counters: - del self.service_context.as_token_counters[name] - - if config.get("backend") not in R.as_token_counters: - logger.warning(f"Token counter backend {config.get('backend')} is not supported.") - continue - try: - config_dict = {k: v for k, v in config.items() if k != "backend"} - self.service_context.as_token_counters[name] = R.as_token_counters[config["backend"]](**config_dict) - logger.info(f"Restarted AS token counter: {name}") - except Exception as e: - logger.error(f"Failed to restart AS token counter '{name}': {e}") - - # llms - if "llms" in restart_config: - llms_config = restart_config["llms"] - assert isinstance(llms_config, dict) - for name, config in llms_config.items(): - if name in self.service_context.llms: - llm = self.service_context.llms.pop(name) - await llm.close() - - if isinstance(config, dict): - config = LLMConfig(**config) - if config.backend not in R.llms: - logger.warning(f"LLM backend {config.backend} is not supported.") - continue - config_dict = config.model_dump(exclude={"backend"}) - config_dict.setdefault("api_key", self.llm_api_key) - config_dict.setdefault("base_url", self.llm_base_url) - self.service_context.llms[name] = R.llms[config.backend](**config_dict) - await self.service_context.llms[name].start() - logger.info(f"Restarted LLM: {name}") - - # embedding_models - if "embedding_models" in restart_config: - embedding_models_config = restart_config["embedding_models"] - assert isinstance(embedding_models_config, dict) - updated_names = set() - for name, config in embedding_models_config.items(): - if name in self.service_context.embedding_models: - embedding_model = self.service_context.embedding_models.pop(name) - await embedding_model.close() - - if isinstance(config, dict): - config = EmbeddingModelConfig(**config) - if config.backend not in R.embedding_models: - logger.warning(f"Embedding model backend {config.backend} is not supported.") - continue - config_dict = config.model_dump(exclude={"backend"}) - config_dict.setdefault("api_key", self.embedding_api_key) - config_dict.setdefault("base_url", self.embedding_base_url) - config_dict.setdefault("cache_dir", working_path / "embedding_cache") - self.service_context.embedding_models[name] = R.embedding_models[config.backend](**config_dict) - await self.service_context.embedding_models[name].start() - logger.info(f"Restarted embedding model: {name}") - updated_names.add(name) - - # update embedding_model attribute for existing vector_stores and file_stores - for name in updated_names: - for vs_name, vs_config in self.service_config.vector_stores.items(): - if vs_config.embedding_model == name and vs_name in self.service_context.vector_stores: - self.service_context.vector_stores[vs_name].embedding_model = ( - self.service_context.embedding_models[name] - ) - logger.info(f"Updated embedding model for vector store: {vs_name}") - for fs_name, fs_config in self.service_config.file_stores.items(): - if fs_config.embedding_model == name and fs_name in self.service_context.file_stores: - self.service_context.file_stores[fs_name].embedding_model = ( - self.service_context.embedding_models[name] - ) - logger.info(f"Updated embedding model for file store: {fs_name}") - - # token_counters - if "token_counters" in restart_config: - token_counters_config = restart_config["token_counters"] - assert isinstance(token_counters_config, dict) - for name, config in token_counters_config.items(): - if name in self.service_context.token_counters: - del self.service_context.token_counters[name] - - if isinstance(config, dict): - config = TokenCounterConfig(**config) - if config.backend not in R.token_counters: - logger.warning(f"Token counter backend {config.backend} is not supported.") - continue - config_dict = config.model_dump(exclude={"backend"}) - self.service_context.token_counters[name] = R.token_counters[config.backend](**config_dict) - logger.info(f"Restarted token counter: {name}") - - # vector_stores - if "vector_stores" in restart_config: - vector_stores_config = restart_config["vector_stores"] - assert isinstance(vector_stores_config, dict) - for name, config in vector_stores_config.items(): - if name in self.service_context.vector_stores: - vector_store = self.service_context.vector_stores.pop(name) - await vector_store.close() - if isinstance(config, dict): - config = VectorStoreConfig(**config) - if config.backend not in R.vector_stores: - logger.warning(f"Vector store backend {config.backend} is not supported.") - continue - config_dict = config.model_dump(exclude={"backend", "embedding_model"}) - config_dict.update( - { - "embedding_model": self.service_context.embedding_models[config.embedding_model], - "db_path": working_path / "vector_store", - }, - ) - self.service_context.vector_stores[name] = R.vector_stores[config.backend](**config_dict) - await self.service_context.vector_stores[name].start() - logger.info(f"Restarted vector store: {name}") - - # file_stores - if "file_stores" in restart_config: - file_stores_config = restart_config["file_stores"] - assert isinstance(file_stores_config, dict) - for name, config in file_stores_config.items(): - if name in self.service_context.file_stores: - file_store = self.service_context.file_stores.pop(name) - await file_store.close() - if isinstance(config, dict): - config = FileStoreConfig(**config) - if config.backend not in R.file_stores: - logger.warning(f"File store backend {config.backend} is not supported.") - continue - config_dict = config.model_dump(exclude={"backend", "embedding_model"}) - config_dict.update( - { - "embedding_model": self.service_context.embedding_models[config.embedding_model], - "db_path": working_path / "file_store", - }, - ) - self.service_context.file_stores[name] = R.file_stores[config.backend](**config_dict) - await self.service_context.file_stores[name].start() - logger.info(f"Restarted file store: {name}") - - # file_watchers - if "file_watchers" in restart_config: - file_watchers_config = restart_config["file_watchers"] - assert isinstance(file_watchers_config, dict) - for name, config in file_watchers_config.items(): - if name in self.service_context.file_watchers: - file_watcher = self.service_context.file_watchers.pop(name) - await file_watcher.close() - if isinstance(config, dict): - config = FileWatcherConfig(**config) - if config.backend not in R.file_watchers: - logger.warning(f"File watcher backend {config.backend} is not supported.") - continue - config_dict = config.model_dump(exclude={"backend", "file_store"}) - config_dict["file_store"] = self.service_context.file_stores[config.file_store] - self.service_context.file_watchers[name] = R.file_watchers[config.backend](**config_dict) - await self.service_context.file_watchers[name].start() - logger.info(f"Restarted file watcher: {name}") - - async def prepare_mcp_servers(self): - """Prepare and initialize MCP server connections.""" - mcp_client = MCPClient(config={"mcpServers": self.service_config.mcp_servers}) - for server_name in self.service_config.mcp_servers.keys(): - try: - tool_calls = await mcp_client.list_tool_calls(server_name=server_name, return_dict=False) - self.service_context.mcp_server_mapping[server_name] = { - tool_call.name: tool_call for tool_call in tool_calls - } - for tool_call in tool_calls: - logger.info(f"list_tool_calls: {server_name}@{tool_call.name} {tool_call.simple_input_dump()}") - except Exception as e: - logger.exception(f"list_tool_calls: {server_name} error: {e}") - - async def close(self) -> bool: - """Close all service components asynchronously.""" - if not self._started: - logger.warning("Application is not started") - return True - - for name, file_watcher in self.service_context.file_watchers.items(): - logger.info(f"Closing file watcher: {name}") - await file_watcher.close() - - for name, file_store in self.service_context.file_stores.items(): - logger.info(f"Closing file store: {name}") - await file_store.close() - - for name, vector_store in self.service_context.vector_stores.items(): - logger.info(f"Closing vector store: {name}") - await vector_store.close() - - for name, llm in self.service_context.llms.items(): - logger.info(f"Closing LLM: {name}") - await llm.close() - - for name, embedding_model in self.service_context.embedding_models.items(): - logger.info(f"Closing embedding model: {name}") - await embedding_model.close() - - self.shutdown_thread_pool() - self.shutdown_ray() - - self._started = False - logger.info("ReMe Application closed") - return False - - def shutdown_thread_pool(self, wait: bool = True): - """Shutdown the thread pool executor.""" - if self.service_context.thread_pool is not None: - self.service_context.thread_pool.shutdown(wait=wait) - - def shutdown_ray(self, wait: bool = True): - """Shutdown Ray cluster if it was initialized.""" - if self.service_config and self.service_config.ray_max_workers > 1: - import ray - - ray.shutdown(_exiting_interpreter=not wait) - - async def __aenter__(self): - """Async context manager entry.""" - return await self.start() - - async def __aexit__(self, exc_type=None, exc_val=None, exc_tb=None): - """Async context manager exit.""" - return await self.close() - - async def execute_flow(self, name: str, **kwargs) -> Response: - """Execute a flow with the given name and parameters.""" - assert name in self.service_context.flows, f"Flow {name} not found" - flow: BaseFlow = self.service_context.flows[name] - return await flow.call(**kwargs) - - async def execute_stream_flow(self, name: str, **kwargs): - """Execute a stream flow with the given name and parameters.""" - assert name in self.service_context.flows, f"Flow {name} not found" - flow: BaseFlow = self.service_context.flows[name] - assert flow.stream is True, "non-stream flow is not supported in execute_stream_flow!" - stream_queue = asyncio.Queue() - task = asyncio.create_task(flow.call(stream_queue=stream_queue, **kwargs)) - async for chunk in execute_stream_task( - stream_queue=stream_queue, - task=task, - task_name=name, - output_format="str", - ): - yield chunk - - @property - def default_llm(self) -> BaseLLM: - """Get the default LLM instance.""" - return self.service_context.llms.get("default") - - def get_llm(self, name: str): - """Get an LLM instance by name.""" - return self.service_context.llms.get(name) - - def update_default_llm_name(self, name: str): - """Update the default LLM name.""" - self.default_llm.model_name = name - - @property - def default_embedding_model(self) -> BaseEmbeddingModel: - """Get the default embedding model instance.""" - return self.service_context.embedding_models.get("default") - - def get_embedding_model(self, name: str): - """Get an embedding model instance by name.""" - return self.service_context.embedding_models.get(name) - - def update_default_embedding_name(self, name: str): - """Update the default embedding model name.""" - self.default_embedding_model.model_name = name - - @property - def default_vector_store(self) -> BaseVectorStore: - """Get the default vector store instance.""" - return self.service_context.vector_stores.get("default") - - def get_vector_store(self, name: str): - """Get a vector store instance by name.""" - return self.service_context.vector_stores.get(name) - - @property - def default_file_store(self) -> BaseFileStore: - """Get the default file store instance.""" - return self.service_context.file_stores.get("default") - - def get_file_store(self, name: str): - """Get a file store instance by name.""" - return self.service_context.file_stores.get(name) - - @property - def default_file_watcher(self) -> BaseFileWatcher: - """Get the default file watcher instance.""" - return self.service_context.file_watchers.get("default") - - def get_file_watcher(self, name: str): - """Get a file watcher instance by name.""" - return self.service_context.file_watchers.get(name) - - @property - def default_token_counter(self) -> BaseTokenCounter: - """Get the default token counter instance.""" - return self.service_context.token_counters.get("default") - - def get_token_counter(self, name: str): - """Get a token counter instance by name.""" - return self.service_context.token_counters.get(name) - - def run_service(self): - """Run the configured service (HTTP, MCP, or CMD).""" - import warnings - - warnings.filterwarnings("ignore", category=DeprecationWarning) - service = R.services[self.service_config.backend](app=self) - service.run() - - async def reset_default_collection(self, collection_name: str): - """Reset the default vector store.""" - await self.service_context.vector_stores["default"].reset_collection(collection_name) diff --git a/reme/core/as_llm/__init__.py b/reme/core/as_llm/__init__.py deleted file mode 100644 index 9cf527af..00000000 --- a/reme/core/as_llm/__init__.py +++ /dev/null @@ -1,9 +0,0 @@ -"""Module for registering AgentScope LLM models.""" - -from agentscope.model import DashScopeChatModel -from agentscope.model import OpenAIChatModel - -from ..registry_factory import R - -R.as_llms.register("openai")(OpenAIChatModel) -R.as_llms.register("dashscope")(DashScopeChatModel) diff --git a/reme/core/as_llm_formatter/__init__.py b/reme/core/as_llm_formatter/__init__.py deleted file mode 100644 index 1c7eee46..00000000 --- a/reme/core/as_llm_formatter/__init__.py +++ /dev/null @@ -1,9 +0,0 @@ -"""Module for registering AgentScope LLM formatters.""" - -from agentscope.formatter import DashScopeChatFormatter - -from .reme_openai_chat_formatter import ReMeOpenAIChatFormatter -from ..registry_factory import R - -R.as_llm_formatters.register("openai")(ReMeOpenAIChatFormatter) -R.as_llm_formatters.register("dashscope")(DashScopeChatFormatter) diff --git a/reme/core/as_llm_formatter/reme_openai_chat_formatter.py b/reme/core/as_llm_formatter/reme_openai_chat_formatter.py deleted file mode 100644 index d6807dfb..00000000 --- a/reme/core/as_llm_formatter/reme_openai_chat_formatter.py +++ /dev/null @@ -1,215 +0,0 @@ -"""ReMeOpenAIChatFormatter""" - -import json -from typing import Any - -from agentscope.formatter import OpenAIChatFormatter -from agentscope.formatter._openai_formatter import ( - _format_openai_image_block, - _to_openai_audio_data, -) -from agentscope.message import Msg, TextBlock, ImageBlock, URLSource -from loguru import logger - - -def _format_openai_video_block(video_block: dict) -> dict[str, Any]: - """Format a video block for OpenAI API. - - Args: - video_block: The video block to format. - - Returns: - A dictionary with video content in OpenAI format. - """ - source = video_block["source"] - if source["type"] == "url": - url = source["url"] - elif source["type"] == "base64": - data = source["data"] - media_type = source["media_type"] - url = f"data:{media_type};base64,{data}" - else: - raise ValueError(f"Unsupported video source type: {source['type']}") - - return { - "type": "video_url", - "video_url": { - "url": url, - }, - } - - -class ReMeOpenAIChatFormatter(OpenAIChatFormatter): - """ReMeOpenAIChatFormatter""" - - async def _format( - self, - msgs: list[Msg], - ) -> list[dict[str, Any]]: - """Format message objects into OpenAI API required format. - - Args: - msgs (`list[Msg]`): - The list of Msg objects to format. - - Returns: - `list[dict[str, Any]]`: - A list of dictionaries, where each dictionary has "name", - "role", and "content" keys. - """ - self.assert_list_of_msgs(msgs) - - messages: list[dict] = [] - i = 0 - while i < len(msgs): - msg = msgs[i] - content_blocks = [] - tool_calls = [] - reasoning_content_blocks = [] - - for block in msg.get_content_blocks(): - typ = block.get("type") - if typ == "text": - content_blocks.append({**block}) - - elif typ == "thinking": - # Collect thinking blocks for reasoning_content field - # This is compatible with models like DeepSeek that support - # extended thinking via reasoning_content field - reasoning_content_blocks.append({**block}) - - elif typ == "tool_use": - tool_calls.append( - { - "id": block.get("id"), - "type": "function", - "function": { - "name": block.get("name"), - "arguments": json.dumps( - block.get("input", {}), - ensure_ascii=False, - ), - }, - }, - ) - - elif typ == "tool_result": - ( - textual_output, - multimodal_data, - ) = self.convert_tool_result_to_string(block["output"]) - - messages.append( - { - "role": "tool", - "tool_call_id": block.get("id"), - "content": (textual_output), # type: ignore[arg-type] - "name": block.get("name"), - }, - ) - - # Then, handle the multimodal data if any - promoted_blocks: list = [] - for url, multimodal_block in multimodal_data: - if multimodal_block["type"] == "image" and self.promote_tool_result_images: - promoted_blocks.extend( - [ - TextBlock( - type="text", - text=f"\n- The image from '{url}': ", - ), - ImageBlock( - type="image", - source=URLSource( - type="url", - url=url, - ), - ), - ], - ) - - if promoted_blocks: - # Insert promoted blocks as new user message(s) - promoted_blocks = [ - TextBlock( - type="text", - text="The following are " - "the image contents from the tool " - f"result of '{block['name']}':", - ), - *promoted_blocks, - TextBlock( - type="text", - text="", - ), - ] - - msgs.insert( - i + 1, - Msg( - name="user", - content=promoted_blocks, - role="user", - ), - ) - - elif typ == "image": - content_blocks.append( - _format_openai_image_block( - block, # type: ignore[arg-type] - ), - ) - - elif typ == "audio": - # Filter out audio content when the multimodal model - # outputs both text and audio, to prevent errors in - # subsequent model calls - if msg.role == "assistant": - continue - input_audio = _to_openai_audio_data(block["source"]) - content_blocks.append( - { - "type": "input_audio", - "input_audio": input_audio, - }, - ) - - elif typ == "video": - # Filter out video content when the multimodal model - # outputs both text and video, to prevent errors in - # subsequent model calls - if msg.role == "assistant": - continue - content_blocks.append( - _format_openai_video_block(block), - ) - - else: - logger.warning( - "Unsupported block type %s in the message, skipped.", - typ, - ) - - msg_openai = { - "role": msg.role, - "name": msg.name, - "content": content_blocks or None, - } - - if tool_calls: - msg_openai["tool_calls"] = tool_calls - - # Add reasoning_content for thinking blocks (compatible with DeepSeek, etc.) - if reasoning_content_blocks: - reasoning_msg = "\n".join(reasoning.get("thinking", "") for reasoning in reasoning_content_blocks) - if reasoning_msg: - msg_openai["reasoning_content"] = reasoning_msg - - # When both content and tool_calls are None, skipped - if msg_openai["content"] or msg_openai.get("tool_calls"): - messages.append(msg_openai) - - # Move to next message - i += 1 - - return messages diff --git a/reme/core/as_token_counter/__init__.py b/reme/core/as_token_counter/__init__.py deleted file mode 100644 index 95910b8e..00000000 --- a/reme/core/as_token_counter/__init__.py +++ /dev/null @@ -1,8 +0,0 @@ -"""Module for registering AgentScope token counters.""" - -from .reme_token_counter import ReMeTokenCounter -from .rule_token_counter import RuleTokenCounter -from ..registry_factory import R - -R.as_token_counters.register("hf")(ReMeTokenCounter) -R.as_token_counters.register("rule")(RuleTokenCounter) diff --git a/reme/core/as_token_counter/reme_token_counter.py b/reme/core/as_token_counter/reme_token_counter.py deleted file mode 100644 index 6722e92f..00000000 --- a/reme/core/as_token_counter/reme_token_counter.py +++ /dev/null @@ -1,123 +0,0 @@ -"""Token counter for ReMe.""" - -import os -from typing import Any - -from agentscope.token import HuggingFaceTokenCounter - -from ..utils import get_logger - -logger = get_logger() - - -class ReMeTokenCounter(HuggingFaceTokenCounter): - """Token counter for CoPaw with configurable tokenizer support. - - This class extends HuggingFaceTokenCounter to provide token counting - functionality with support for both local and remote tokenizers, - as well as HuggingFace mirror for users in China. - - Attributes: - pretrained_model_name_or_path: The tokenizer model path or "default" for local tokenizer. - use_mirror: Whether to use HuggingFace mirror. - token_count_estimate_divisor: Divisor for token estimation. - """ - - def __init__( - self, - pretrained_model_name_or_path: str, - use_mirror: bool = True, - token_count_estimate_divisor: float = 3.75, - **kwargs, - ): - """Initialize the token counter with the specified configuration. - - Args: - pretrained_model_name_or_path: The tokenizer model path. - use_mirror: Whether to use the HuggingFace mirror - (https://hf-mirror.com) for downloading tokenizers. Useful for - users in China. - token_count_estimate_divisor: Divisor for estimating tokens when - tokenizer is unavailable. Defaults to 3.75. - **kwargs: Additional keyword arguments passed to HuggingFaceTokenCounter. - """ - self.pretrained_model_name_or_path = pretrained_model_name_or_path - self.use_mirror = use_mirror - self.token_count_estimate_divisor = token_count_estimate_divisor - - # Set HuggingFace endpoint for mirror support - if use_mirror: - mirror = "https://hf-mirror.com" - else: - mirror = "https://huggingface.co" - - os.environ["HF_ENDPOINT"] = mirror - - # if the huggingface is already imported in other dependencies, - # we need to set the endpoint manually - import huggingface_hub.constants - - huggingface_hub.constants.ENDPOINT = mirror - huggingface_hub.constants.HUGGINGFACE_CO_URL_TEMPLATE = mirror + "/{repo_id}/resolve/{revision}/{filename}" - - try: - super().__init__( - pretrained_model_name_or_path=self.pretrained_model_name_or_path, - use_mirror=use_mirror, - use_fast=True, - trust_remote_code=True, - **kwargs, - ) - self._tokenizer_available = True - - except Exception as e: - logger.error(f"Failed to initialize tokenizer {e}") - self._tokenizer_available = False - - async def count( - self, - messages: list[dict], - tools: list[dict] | None = None, - text: str | None = None, - **kwargs: Any, - ) -> int: - """Count tokens in messages or text. - - If text is provided, counts tokens directly in the text string. - Otherwise, counts tokens in the messages using the parent class method. - - Args: - messages: List of message dictionaries in chat format. - tools: Optional list of tool definitions for token counting. - text: Optional text string to count tokens directly. - **kwargs: Additional keyword arguments passed to parent count method. - - Returns: - The number of tokens, guaranteed to be at least the estimated minimum. - """ - if text: - if self._tokenizer_available: - try: - token_ids = self.tokenizer.encode(text) - return max(len(token_ids), self.estimate_tokens(text)) - except Exception as e: - logger.exception("Failed to encode text with tokenizer: %s", e) - return self.estimate_tokens(text) - else: - return self.estimate_tokens(text) - else: - return await super().count(messages, tools, **kwargs) - - def estimate_tokens(self, text: str) -> int: - """Estimate the number of tokens in a text string. - - Provides a fast character-based estimation as a fallback or lower bound. - Uses the configured divisor from instance settings. - - Args: - text: The text string to estimate tokens for. - - Returns: - The estimated number of tokens in the text string. - """ - return int(len(text.encode("utf-8")) / self.token_count_estimate_divisor + 0.5) diff --git a/reme/core/as_token_counter/rule_token_counter.py b/reme/core/as_token_counter/rule_token_counter.py deleted file mode 100644 index 1201e185..00000000 --- a/reme/core/as_token_counter/rule_token_counter.py +++ /dev/null @@ -1,78 +0,0 @@ -"""Rule-based token counter for fast estimation without loading tokenizer.""" - -from typing import Any - -from agentscope.token import HuggingFaceTokenCounter - - -class RuleTokenCounter(HuggingFaceTokenCounter): - """Lightweight token counter using rule-based estimation only. - - This class provides fast token estimation without loading any tokenizer, - useful when exact token counts are not critical or for quick approximations. - - Attributes: - token_count_estimate_divisor: Divisor for token estimation. - """ - - def __init__( - self, - token_count_estimate_divisor: float = 3.75, - **_kwargs, - ): - """Initialize the rule-based token counter. - - Args: - token_count_estimate_divisor: Divisor for estimating tokens. - Defaults to 3.75 (approximately 4 characters per token). - **kwargs: Additional keyword arguments (ignored). - """ - self.token_count_estimate_divisor = token_count_estimate_divisor - # Skip tokenizer initialization from parent - self._tokenizer_available = False - - async def count( - self, - messages: list[dict], - _tools: list[dict] | None = None, - text: str | None = None, - **_kwargs: Any, - ) -> int: - """Count tokens using rule-based estimation. - - Args: - messages: List of message dictionaries in chat format. - _tools: Optional list of tool definitions (ignored). - text: Optional text string to count tokens directly. - **_kwargs: Additional keyword arguments (ignored). - - Returns: - The estimated number of tokens. - """ - if text: - return self.estimate_tokens(text) - - # Estimate from messages - total_text = "" - for msg in messages: - content = msg.get("content", "") - if isinstance(content, str): - total_text += content - elif isinstance(content, list): - for part in content: - if isinstance(part, dict) and "text" in part: - total_text += part["text"] - return self.estimate_tokens(total_text) - - def estimate_tokens(self, text: str) -> int: - """Estimate the number of tokens in a text string. - - Uses character-based estimation with the configured divisor. - - Args: - text: The text string to estimate tokens for. - - Returns: - The estimated number of tokens in the text string. - """ - return int(len(text.encode("utf-8")) / self.token_count_estimate_divisor + 0.5) diff --git a/reme/core/base_dict.py b/reme/core/base_dict.py deleted file mode 100644 index b96b4c24..00000000 --- a/reme/core/base_dict.py +++ /dev/null @@ -1,41 +0,0 @@ -"""Module providing a dictionary subclass with attribute-style access and pickling support.""" - -from typing import Generic, TypeVar - -_KT = TypeVar("_KT") -_VT = TypeVar("_VT") - - -class BaseDict(dict, Generic[_KT, _VT]): - """A dictionary subclass that enables accessing and modifying keys as attributes.""" - - def __getattr__(self, name: str) -> _VT: - """Retrieve a dictionary item as an attribute.""" - try: - return self[name] - except KeyError as e: - raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e - - def __setattr__(self, name: str, value: _VT) -> None: - """Assign a value to a dictionary item using attribute syntax.""" - self[name] = value - - def __delattr__(self, name: str) -> None: - """Remove a dictionary item using attribute syntax.""" - try: - # Delete item from dict via key - del self[name] - except KeyError as e: - raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e - - def __getstate__(self) -> dict: - """Return the dictionary representation for pickling.""" - return dict(self) - - def __setstate__(self, state: dict) -> None: - """Restore the dictionary state from a pickled object.""" - self.update(state) - - def __reduce__(self): - """Define the reconstruction logic for pickling processes.""" - return self.__class__, (), self.__getstate__() diff --git a/reme/core/embedding/__init__.py b/reme/core/embedding/__init__.py deleted file mode 100644 index 6431703e..00000000 --- a/reme/core/embedding/__init__.py +++ /dev/null @@ -1,15 +0,0 @@ -"""embedding""" - -from .base_embedding_model import BaseEmbeddingModel -from .openai_embedding_model import OpenAIEmbeddingModel -from .openai_embedding_model_sync import OpenAIEmbeddingModelSync -from ..registry_factory import R - -__all__ = [ - "BaseEmbeddingModel", - "OpenAIEmbeddingModel", - "OpenAIEmbeddingModelSync", -] - -R.embedding_models.register("openai")(OpenAIEmbeddingModel) -R.embedding_models.register("openai_sync")(OpenAIEmbeddingModelSync) diff --git a/reme/core/embedding/base_embedding_model.py b/reme/core/embedding/base_embedding_model.py deleted file mode 100644 index e40526d0..00000000 --- a/reme/core/embedding/base_embedding_model.py +++ /dev/null @@ -1,599 +0,0 @@ -"""Base embedding model interface for ReMe. - -Defines the abstract base class and standard API for all embedding model implementations. -""" - -import asyncio -import hashlib -import json -import time -from abc import ABC -from collections import OrderedDict -from pathlib import Path - -from loguru import logger - -from ..schema import VectorNode, MemoryChunk - - -class BaseEmbeddingModel(ABC): - """Abstract base class for embedding model implementations. - - Provides a standard interface for text-to-vector generation with - built-in batching, retry logic, and error handling. - """ - - def __init__( - self, - api_key: str | None = None, - base_url: str | None = None, - model_name: str = "", - dimensions: int = 1024, - use_dimensions: bool = False, - max_batch_size: int = 10, - max_retries: int = 3, - raise_exception: bool = True, - max_input_length: int = 8192, - cache_dir: str | Path = ".reme", - max_cache_size: int = 2000, - enable_cache: bool = True, - **kwargs, - ): - """Initialize model configuration and parameters. - - Args: - api_key: API key for the embedding service - base_url: Base URL for the embedding service - model_name: Name of the embedding model - dimensions: Vector dimensions of the embeddings - use_dimensions: Whether to pass dimensions parameter to API (some APIs don't support it) - max_batch_size: Maximum batch size for embedding requests - max_retries: Maximum number of retry attempts on failure - raise_exception: Whether to raise exceptions on failure - max_input_length: Maximum input text length - max_cache_size: Maximum number of embeddings to cache in memory (LRU) - enable_cache: Whether to enable embedding cache - **kwargs: Additional model-specific parameters - """ - self.api_key: str | None = api_key - self.base_url: str | None = base_url - self.model_name = model_name - self.dimensions = dimensions - self.use_dimensions = use_dimensions - self.max_batch_size = max_batch_size - self.max_retries = max_retries - self.raise_exception = raise_exception - self.max_input_length = max_input_length - self.cache_dir = cache_dir - self.max_cache_size = max_cache_size - self.enable_cache = enable_cache - self.kwargs = kwargs - - # Initialize LRU cache for embeddings - self._embedding_cache: OrderedDict[str, list[float]] = OrderedDict() - self._cache_hits = 0 - self._cache_misses = 0 - - self.cache_path: Path = Path(self.cache_dir) - self.cache_path.mkdir(parents=True, exist_ok=True) - - def _truncate_text(self, text: str) -> str: - """Truncate text to max_input_length if it exceeds the limit.""" - if len(text) > self.max_input_length: - logger.warning(f"Text length {len(text)} exceeds {self.max_input_length}, truncating") - return text[: self.max_input_length] - return text - - def _truncate_texts(self, texts: list[str]) -> list[str]: - """Truncate a list of texts to max_input_length.""" - return [self._truncate_text(text) for text in texts] - - def _validate_and_adjust_embedding(self, embedding: list[float]) -> list[float]: - """Validate and adjust embedding dimensions to match expected dimensions. - - Args: - embedding: The embedding vector to validate - - Returns: - Embedding vector adjusted to match self.dimensions - """ - actual_len = len(embedding) - if actual_len == self.dimensions: - return embedding - - elif actual_len < self.dimensions: - logger.warning( - f"[ACTUAL_EMB_LENGTH]Embedding dimensions {actual_len} is less than expected {self.dimensions}, " - f"padding with zeros", - ) - return embedding + [0.0] * (self.dimensions - actual_len) - - else: - logger.warning( - f"[ACTUAL_EMB_LENGTH]Embedding dimensions {actual_len} is greater than expected {self.dimensions}, " - f"truncating to {self.dimensions}", - ) - return embedding[: self.dimensions] - - def _get_cache_key(self, text: str, dimensions: int) -> str: - """Generate a cache key by hashing text + model_name + dimensions. - - This ensures that the same text produces different cache keys when - using different models or dimensions. - - Args: - text: Input text to hash - dimensions: Vector dimensions of the embeddings - - Returns: - SHA256 hash combining text, model name, and dimensions - """ - # Combine text, model_name, and dimensions to create unique cache key - cache_string = f"{text}|{self.model_name}|{dimensions}" - return hashlib.sha256(cache_string.encode("utf-8")).hexdigest() - - def _get_cache_file_path(self) -> Path: - """Get the path to the cache file. - - Returns: - Path to the embedding cache JSONL file - """ - return self.cache_path / "embedding_cache.jsonl" - - def _load_cache(self) -> None: - """Load embedding cache from disk (JSONL format). - - Each line in the JSONL file contains a JSON object with: - - key: the cache key (SHA256 hash) - - embedding: the embedding vector (list of floats) - - Loads in reverse order (newest first) to prioritize recent embeddings - when max_cache_size is smaller than the file content. - """ - if not self.enable_cache: - return - - cache_file = self._get_cache_file_path() - if not cache_file.exists(): - logger.info(f"No cache file found at {cache_file}, starting with empty cache") - return - - try: - load_start = time.time() - # Read all lines first (to load in reverse order) - with open(cache_file, "r", encoding="utf-8") as f: - lines = f.readlines() - - loaded_count = 0 - # Load in reverse order (newest entries first) - for _, line in enumerate(reversed(lines), 1): - line = line.strip() - if not line: - continue - try: - data = json.loads(line) - if not data: - continue - # Each line is {cache_key: embedding} - cache_key, embedding = next(iter(data.items())) - - if cache_key and embedding and isinstance(embedding, list): - # Skip if already loaded (keep the newest) - if cache_key in self._embedding_cache: - continue - - if len(embedding) != self.dimensions: - logger.warning( - f"Embedding dimensions mismatch for cache key {cache_key}, " - f"expected {self.dimensions}, got {len(embedding)}", - ) - continue - - # Respect max_cache_size during loading - if len(self._embedding_cache) >= self.max_cache_size: - logger.info( - f"Cache size limit reached ({self.max_cache_size}), " - f"loaded {loaded_count} newest entries", - ) - break - self._embedding_cache[cache_key] = embedding - loaded_count += 1 - except json.JSONDecodeError as e: - logger.warning(f"Failed to parse line in cache file: {e}") - continue - - logger.info( - f"Loaded {loaded_count} embeddings from cache file: {cache_file} in {time.time() - load_start:.2f}s", - ) - except Exception as e: - logger.error(f"Failed to load cache from {cache_file}: {e}, deleting cache file") - try: - cache_file.unlink() - logger.info(f"Deleted corrupted cache file: {cache_file}") - except Exception as del_e: - logger.error(f"Failed to delete cache file {cache_file}: {del_e}") - - def _save_cache(self) -> None: - """Save embedding cache to disk (JSONL format). - - Each line contains a JSON object with the cache key and embedding vector. - Only saves if cache is non-empty. - """ - if not self.enable_cache: - return - - logger.info(f"Attempting to save cache, current size: {len(self._embedding_cache)}") - if not self._embedding_cache: - logger.info("Cache is empty, skipping save") - return - - cache_file = self._get_cache_file_path() - try: - with open(cache_file, "w", encoding="utf-8") as f: - for cache_key, embedding in self._embedding_cache.items(): - if len(embedding) != self.dimensions: - logger.warning( - f"Embedding dimensions mismatch for cache key {cache_key}, " - f"expected {self.dimensions}, got {len(embedding)}", - ) - continue - cache_entry = {cache_key: embedding} - f.write(json.dumps(cache_entry, ensure_ascii=False) + "\n") - - logger.info(f"Saved {len(self._embedding_cache)} embeddings to cache file: {cache_file}") - except Exception as e: - logger.error(f"Failed to save cache to {cache_file}: {e}") - - def _get_from_cache(self, text: str) -> list[float] | None: - """Retrieve embedding from cache if it exists. - - Args: - text: Input text to look up - - Returns: - Cached embedding vector or None if not found - """ - if not self.enable_cache: - return None - - cache_key = self._get_cache_key(text, self.dimensions) - if cache_key in self._embedding_cache: - embeddings: list[float] = self._embedding_cache[cache_key] - - # Validate embedding dimensions match expected dimensions - if len(embeddings) != self.dimensions: - logger.warning( - f"Cached embedding dimensions mismatch: expected {self.dimensions}, " - f"got {len(embeddings)}. Removing invalid cache entry.", - ) - del self._embedding_cache[cache_key] - self._cache_misses += 1 - return None - - # Move to end (most recently used) - self._embedding_cache.move_to_end(cache_key) - self._cache_hits += 1 - text_preview = text[:50] + "..." if len(text) > 50 else text - logger.info(f"Cache hit for text: {text_preview} (hits: {self._cache_hits}, misses: {self._cache_misses})") - return embeddings - - self._cache_misses += 1 - return None - - def _put_to_cache(self, text: str, embedding: list[float]) -> None: - """Store embedding in cache with LRU eviction. - - Args: - text: Input text used as cache key - embedding: Embedding vector to cache - """ - if not self.enable_cache: - return - - if self.max_cache_size <= 0: - return - - cache_key = self._get_cache_key(text, self.dimensions) - if len(embedding) != self.dimensions: - logger.warning( - f"[PUT_TO_CACHE] Embedding dimensions mismatch for cache key {cache_key}, " - f"expected {self.dimensions}, got real length {len(embedding)}", - ) - return - - # Remove the oldest entry if cache is full - if len(self._embedding_cache) >= self.max_cache_size and cache_key not in self._embedding_cache: - self._embedding_cache.popitem(last=False) - - self._embedding_cache[cache_key] = embedding - self._embedding_cache.move_to_end(cache_key) - - def get_cache_stats(self) -> dict[str, int]: - """Get cache statistics. - - Returns: - Dictionary with cache size, hits, misses, and hit rate - """ - total_requests = self._cache_hits + self._cache_misses - hit_rate = self._cache_hits / total_requests if total_requests > 0 else 0.0 - return { - "cache_size": len(self._embedding_cache), - "max_cache_size": self.max_cache_size, - "cache_hits": self._cache_hits, - "cache_misses": self._cache_misses, - "hit_rate": hit_rate, - } - - def clear_cache(self) -> None: - """Clear the embedding cache and reset statistics.""" - self._embedding_cache.clear() - self._cache_hits = 0 - self._cache_misses = 0 - - async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]: - """Internal async implementation for calling the embedding API with batch input.""" - - def _get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]: - """Internal synchronous implementation for calling the embedding API with batch input.""" - - async def get_embedding(self, input_text: str, **kwargs) -> list[float]: - """Async get embedding for a single text with exponential backoff retries.""" - truncated_text = self._truncate_text(input_text) - - # Check cache first - cached_embedding = self._get_from_cache(truncated_text) - if cached_embedding is not None: - return cached_embedding - - # Cache miss - compute embedding - for i in range(self.max_retries): - try: - result = await self._get_embeddings([truncated_text], **kwargs) - embedding = self._validate_and_adjust_embedding(result[0]) - # Store in cache - self._put_to_cache(truncated_text, embedding) - return embedding - except Exception as e: - logger.error(f"Model {self.model_name} failed: {e}") - if i == self.max_retries - 1: - if self.raise_exception: - raise - return [] - await asyncio.sleep(i + 1) - return [] - - async def get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]: - """Async get embeddings with automatic batching and exponential backoff retries.""" - # Truncate all input texts first - truncated_texts = self._truncate_texts(input_text) - - # Check cache for each text and separate cached vs uncached - results: list[list[float] | None] = [None] * len(truncated_texts) - texts_to_compute: list[tuple[int, str]] = [] # (original_index, text) - - for idx, text in enumerate(truncated_texts): - cached = self._get_from_cache(text) - if cached is not None: - results[idx] = cached - else: - texts_to_compute.append((idx, text)) - - # If all texts were cached, return early - if not texts_to_compute: - return [r for r in results if r is not None] - - # Compute embeddings for uncached texts in batches - uncached_texts = [text for _, text in texts_to_compute] - for i in range(0, len(uncached_texts), self.max_batch_size): - batch_texts = uncached_texts[i : i + self.max_batch_size] - batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]] - - # Process each batch with retry logic - for retry in range(self.max_retries): - try: - batch_embeddings = await self._get_embeddings(batch_texts, **kwargs) - if batch_embeddings: - # Store results and cache them - for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings): - adjusted_embedding = self._validate_and_adjust_embedding(embedding) - results[orig_idx] = adjusted_embedding - self._put_to_cache(text, adjusted_embedding) - break - except Exception as e: - logger.error(f"Model {self.model_name} batch failed: {e}") - if retry == self.max_retries - 1: - if self.raise_exception: - raise - else: - await asyncio.sleep(retry + 1) - - return [r for r in results if r is not None] - - def get_embedding_sync(self, input_text: str, **kwargs) -> list[float]: - """Synchronous get embedding for a single text with retry logic.""" - truncated_text = self._truncate_text(input_text) - - # Check cache first - cached_embedding = self._get_from_cache(truncated_text) - if cached_embedding is not None: - return cached_embedding - - # Cache miss - compute embedding - for i in range(self.max_retries): - try: - result = self._get_embeddings_sync([truncated_text], **kwargs) - embedding = self._validate_and_adjust_embedding(result[0]) - # Store in cache - self._put_to_cache(truncated_text, embedding) - return embedding - except Exception as exc: - logger.error(f"Model {self.model_name} failed: {exc}") - if i == self.max_retries - 1: - if self.raise_exception: - raise - return [] - time.sleep(i + 1) - return [] - - def get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]: - """Synchronous get embeddings with automatic batching and retry logic.""" - # Truncate all input texts first - truncated_texts = self._truncate_texts(input_text) - - # Check cache for each text and separate cached vs uncached - results: list[list[float] | None] = [None] * len(truncated_texts) - texts_to_compute: list[tuple[int, str]] = [] # (original_index, text) - - for idx, text in enumerate(truncated_texts): - cached = self._get_from_cache(text) - if cached is not None: - results[idx] = cached - else: - texts_to_compute.append((idx, text)) - - # If all texts were cached, return early - if not texts_to_compute: - return [r for r in results if r is not None] - - # Compute embeddings for uncached texts in batches - uncached_texts = [text for _, text in texts_to_compute] - for i in range(0, len(uncached_texts), self.max_batch_size): - batch_texts = uncached_texts[i : i + self.max_batch_size] - batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]] - - # Process each batch with retry logic - for retry in range(self.max_retries): - try: - batch_embeddings = self._get_embeddings_sync(batch_texts, **kwargs) - if batch_embeddings: - # Store results and cache them - for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings): - adjusted_embedding = self._validate_and_adjust_embedding(embedding) - results[orig_idx] = adjusted_embedding - self._put_to_cache(text, adjusted_embedding) - break - except Exception as exc: - logger.error(f"Model {self.model_name} batch failed: {exc}") - if retry == self.max_retries - 1: - if self.raise_exception: - raise - else: - time.sleep(retry + 1) - - return [r for r in results if r is not None] - - async def get_node_embedding(self, node: VectorNode, **kwargs) -> VectorNode: - """Async generate and populate vector field for a single VectorNode object.""" - node.vector = await self.get_embedding(node.content, **kwargs) - return node - - async def get_node_embeddings(self, nodes: list[VectorNode], **kwargs) -> list[VectorNode]: - """Async generate and populate vector fields for a batch of VectorNode objects.""" - contents = [node.content for node in nodes] - embeddings: list[list[float]] = await self.get_embeddings(contents, **kwargs) - - if len(embeddings) == len(nodes): - for node, vec in zip(nodes, embeddings): - node.vector = vec - else: - logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(nodes)} nodes") - return nodes - - def get_node_embedding_sync(self, node: VectorNode, **kwargs) -> VectorNode: - """Synchronously generate and populate vector field for a single VectorNode object.""" - node.vector = self.get_embedding_sync(node.content, **kwargs) - return node - - def get_node_embeddings_sync(self, nodes: list[VectorNode], **kwargs) -> list[VectorNode]: - """Synchronously generate and populate vector fields for a batch of VectorNode objects.""" - contents = [node.content for node in nodes] - embeddings: list[list[float]] = self.get_embeddings_sync(contents, **kwargs) - - if len(embeddings) == len(nodes): - for node, vec in zip(nodes, embeddings): - node.vector = vec - else: - logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(nodes)} nodes") - return nodes - - async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk: - """Async generate and populate embedding field for a single MemoryChunk object. - - Args: - chunk: MemoryChunk object containing text to embed - **kwargs: Additional arguments passed to the embedding model - - Returns: - The same MemoryChunk object with populated embedding field - """ - chunk.embedding = await self.get_embedding(chunk.text, **kwargs) - return chunk - - async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]: - """Async generate and populate embedding fields for a batch of MemoryChunk objects. - - Args: - chunks: List of MemoryChunk objects containing text to embed - **kwargs: Additional arguments passed to the embedding model - - Returns: - The same list of MemoryChunk objects with populated embedding fields - """ - texts = [chunk.text for chunk in chunks] - embeddings: list[list[float]] = await self.get_embeddings(texts, **kwargs) - - if len(embeddings) == len(chunks): - for chunk, vec in zip(chunks, embeddings): - chunk.embedding = vec - else: - logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(chunks)} chunks") - return chunks - - def get_chunk_embedding_sync(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk: - """Synchronously generate and populate embedding field for a single MemoryChunk object. - - Args: - chunk: MemoryChunk object containing text to embed - **kwargs: Additional arguments passed to the embedding model - - Returns: - The same MemoryChunk object with populated embedding field - """ - chunk.embedding = self.get_embedding_sync(chunk.text, **kwargs) - return chunk - - def get_chunk_embeddings_sync(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]: - """Synchronously generate embeddings for a batch of MemoryChunk objects. - - Args: - chunks: List of MemoryChunk objects containing text to embed - **kwargs: Additional arguments passed to the embedding model - - Returns: - The same list of MemoryChunk objects with populated embedding fields - """ - texts = [chunk.text for chunk in chunks] - embeddings: list[list[float]] = self.get_embeddings_sync(texts, **kwargs) - - if len(embeddings) == len(chunks): - for chunk, vec in zip(chunks, embeddings): - chunk.embedding = vec - else: - logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(chunks)} chunks") - return chunks - - def start_sync(self): - """Synchronously initialize resources and load cache.""" - self._load_cache() - - async def start(self): - """Asynchronously initialize resources and load cache.""" - self._load_cache() - - def close_sync(self): - """Synchronously release resources and close connections.""" - self._save_cache() - - async def close(self): - """Asynchronously release resources and close connections.""" - self._save_cache() diff --git a/reme/core/embedding/openai_embedding_model.py b/reme/core/embedding/openai_embedding_model.py deleted file mode 100644 index d686c83c..00000000 --- a/reme/core/embedding/openai_embedding_model.py +++ /dev/null @@ -1,62 +0,0 @@ -"""Asynchronous OpenAI-compatible embedding model implementation for ReMe.""" - -from typing import Literal - -from openai import AsyncOpenAI - -from .base_embedding_model import BaseEmbeddingModel - - -class OpenAIEmbeddingModel(BaseEmbeddingModel): - """Asynchronous embedding model implementation compatible with OpenAI-style APIs.""" - - def __init__(self, encoding_format: Literal["float", "base64"] = "float", **kwargs): - """Initialize the OpenAI async embedding model with API credentials and configuration.""" - super().__init__(**kwargs) - self.encoding_format: Literal["float", "base64"] = encoding_format - - # Lazy client initialization - self._client = None - - def _create_client(self): - """Create and return an internal AsyncOpenAI client instance.""" - return AsyncOpenAI(api_key=self.api_key, base_url=self.base_url) - - @property - def client(self): - """Lazily create and return the AsyncOpenAI client.""" - if self._client is None: - self._client = self._create_client() - return self._client - - async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]: - """Fetch embeddings from the API for a batch of strings.""" - create_kwargs: dict = { - "model": self.model_name, - "input": input_text, - "encoding_format": self.encoding_format, - **self.kwargs, - **kwargs, - } - if self.use_dimensions: - create_kwargs["dimensions"] = self.dimensions - - completion = await self.client.embeddings.create(**create_kwargs) - - result_emb = [[] for _ in range(len(input_text))] - for emb in completion.data: - # BGE-M3 returns dense_embedding instead of embedding; use as fallback - vec = getattr(emb, "embedding", None) or getattr(emb, "dense_embedding", None) - result_emb[emb.index] = list(vec) if vec is not None else [] - return result_emb - - async def start(self): - """Initialize the asynchronous OpenAI embedding model and load cache.""" - await super().start() - - async def close(self): - """Close the asynchronous OpenAI client and release network resources.""" - if self._client is not None: - await self._client.close() - self._client = None - await super().close() diff --git a/reme/core/embedding/openai_embedding_model_sync.py b/reme/core/embedding/openai_embedding_model_sync.py deleted file mode 100644 index fc05de9c..00000000 --- a/reme/core/embedding/openai_embedding_model_sync.py +++ /dev/null @@ -1,39 +0,0 @@ -"""Synchronous OpenAI-compatible embedding model implementation for ReMe.""" - -from openai import OpenAI - -from .openai_embedding_model import OpenAIEmbeddingModel - - -class OpenAIEmbeddingModelSync(OpenAIEmbeddingModel): - """Synchronous embedding model implementation that extends the asynchronous OpenAI model.""" - - def _create_client(self): - """Create and return an internal synchronous OpenAI client instance.""" - return OpenAI(api_key=self.api_key, base_url=self.base_url) - - def _get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]: - """Fetch embeddings synchronously from the API for a batch of strings.""" - create_kwargs: dict = { - "model": self.model_name, - "input": input_text, - "encoding_format": self.encoding_format, - **self.kwargs, - **kwargs, - } - if self.use_dimensions: - create_kwargs["dimensions"] = self.dimensions - - completion = self.client.embeddings.create(**create_kwargs) - - result_emb = [[] for _ in range(len(input_text))] - for emb in completion.data: - result_emb[emb.index] = emb.embedding - return result_emb - - def close_sync(self): - """Close the synchronous OpenAI client and release network resources.""" - if self._client is not None: - self._client.close() - self._client = None - super().close_sync() diff --git a/reme/core/enumeration/__init__.py b/reme/core/enumeration/__init__.py deleted file mode 100644 index c709c20d..00000000 --- a/reme/core/enumeration/__init__.py +++ /dev/null @@ -1,19 +0,0 @@ -"""enumeration""" - -from .chunk_enum import ChunkEnum -from .http_enum import HttpEnum -from .json_schema_enum import JsonSchemaEnum -from .memory_source import MemorySource -from .memory_type import MemoryType -from .registry_enum import RegistryEnum -from .role import Role - -__all__ = [ - "ChunkEnum", - "HttpEnum", - "JsonSchemaEnum", - "MemorySource", - "MemoryType", - "RegistryEnum", - "Role", -] diff --git a/reme/core/enumeration/chunk_enum.py b/reme/core/enumeration/chunk_enum.py deleted file mode 100644 index 736d9c67..00000000 --- a/reme/core/enumeration/chunk_enum.py +++ /dev/null @@ -1,31 +0,0 @@ -"""Defines the types of data chunks used in streaming responses.""" - -from enum import Enum - - -class ChunkEnum(str, Enum): - """Enumeration of possible chunk categories for stream processing.""" - - # Internal reasoning or chain-of-thought process - THINK = "think" - - # The final generated response content - ANSWER = "answer" - - # Metadata or calls related to external tools - TOOL = "tool" - - # Resource consumption and token usage statistics - USAGE = "usage" - - # Error messages or exception details - ERROR = "error" - - # Signal indicating the start of a new ReAct step - STEP_START = "step_start" - - # Tool execution result - TOOL_RESULT = "tool_result" - - # Final signal indicating the completion of the stream - DONE = "done" diff --git a/reme/core/enumeration/http_enum.py b/reme/core/enumeration/http_enum.py deleted file mode 100644 index 19622242..00000000 --- a/reme/core/enumeration/http_enum.py +++ /dev/null @@ -1,22 +0,0 @@ -"""Provides a collection of standard HTTP request methods.""" - -from enum import Enum - - -class HttpEnum(str, Enum): - """Enumeration of supported HTTP methods for network requests.""" - - # Retrieves data from a specified resource - GET = "get" - - # Submits data to be processed to a specified resource - POST = "post" - - # Identical to GET but only retrieves the response headers - HEAD = "head" - - # Uploads or replaces the representation of a target resource - PUT = "put" - - # Deletes the specified resource from the server - DELETE = "delete" diff --git a/reme/core/enumeration/json_schema_enum.py b/reme/core/enumeration/json_schema_enum.py deleted file mode 100644 index d66882e2..00000000 --- a/reme/core/enumeration/json_schema_enum.py +++ /dev/null @@ -1,38 +0,0 @@ -"""Defines the standard data types supported by JSON Schema. - -This enum maps common JSON Schema primitive types to their corresponding -Python runtime types, and provides a convenient string representation -compatible with JSON Schema (`"string"`, `"number"`, etc.). -""" - -from enum import Enum - - -class JsonSchemaEnum(Enum): - """Enumeration of valid JSON Schema data types. - - The enum value is the corresponding Python type, while the string - representation (`str(...)`) is the canonical JSON Schema type name. - """ - - # Textual data - STRING = str - - # Numeric values, including integers and floats - NUMBER = float - - # Integer-only numeric values - INTEGER = int - - # JSON objects (key-value mappings) - OBJECT = dict - - # Ordered JSON lists/arrays - ARRAY = list - - # Boolean values: true / false - BOOLEAN = bool - - def __str__(self) -> str: - """Return the lowercase JSON Schema type name for this enum member.""" - return self.name.lower() diff --git a/reme/core/enumeration/memory_source.py b/reme/core/enumeration/memory_source.py deleted file mode 100644 index 05174549..00000000 --- a/reme/core/enumeration/memory_source.py +++ /dev/null @@ -1,11 +0,0 @@ -"""Memory source types.""" - -from enum import Enum - - -class MemorySource(str, Enum): - """Source of memory data.""" - - MEMORY = "memory" - - SESSIONS = "sessions" diff --git a/reme/core/enumeration/memory_type.py b/reme/core/enumeration/memory_type.py deleted file mode 100644 index b9f5ed29..00000000 --- a/reme/core/enumeration/memory_type.py +++ /dev/null @@ -1,33 +0,0 @@ -"""Defines the high-level categories of memory managed by ReMe. - -This enumeration is used across the system to tag, route, and store different -kinds of memories (identity, personal context, procedures, tools, etc.). -""" - -from enum import Enum - - -class MemoryType(str, Enum): - """Enumeration of memory categories used by the memory subsystem. - - These types describe *what* a piece of memory is about, which guides - storage, retrieval, and summarization strategies. - """ - - # Long‑term, relatively stable attributes about the user (name, roles, etc.) - IDENTITY = "identity" - - # User-specific preferences, habits, and evolving personal context - PERSONAL = "personal" - - # How‑to knowledge, workflows, and step‑by‑step instructions - PROCEDURAL = "procedural" - - # Information learned about tools, APIs, and their usage patterns - TOOL = "tool" - - # Condensed representation of larger memory collections - SUMMARY = "summary" - - # Raw chronological interaction history, typically before summarization - HISTORY = "history" diff --git a/reme/core/enumeration/registry_enum.py b/reme/core/enumeration/registry_enum.py deleted file mode 100644 index 68450e6a..00000000 --- a/reme/core/enumeration/registry_enum.py +++ /dev/null @@ -1,31 +0,0 @@ -"""Defines the registry categories for core components of the system.""" - -from enum import Enum - - -class RegistryEnum(str, Enum): - """Enumeration of component types registered within the application lifecycle.""" - - # Large Language Model interfaces - LLM = "llm" - - # Models used for generating vector embeddings - EMBEDDING_MODEL = "embedding_model" - - # Databases or storage systems for vector search - VECTOR_STORE = "vector_store" - - # Databases or storage systems for long-term file storage - FILE_STORE = "file_store" - - # Atomic operations or functional units - OP = "op" - - # Orchestrated sequences of operations or workflows - FLOW = "flow" - - # External APIs or shared internal services - SERVICE = "service" - - # Utilities for tracking and limiting token consumption - TOKEN_COUNTER = "token_counter" diff --git a/reme/core/enumeration/role.py b/reme/core/enumeration/role.py deleted file mode 100644 index 4acad7e5..00000000 --- a/reme/core/enumeration/role.py +++ /dev/null @@ -1,19 +0,0 @@ -"""Defines the participant roles in a chat completion sequence.""" - -from enum import Enum - - -class Role(str, Enum): - """Enumeration of standard personas involved in a conversation flow.""" - - # High-level instructions to guide the model's behavior - SYSTEM = "system" - - # Input or queries provided by the human user - USER = "user" - - # Responses or messages generated by the AI model - ASSISTANT = "assistant" - - # Output or results returned from external tool executions - TOOL = "tool" diff --git a/reme/core/file_store/__init__.py b/reme/core/file_store/__init__.py deleted file mode 100644 index f071956d..00000000 --- a/reme/core/file_store/__init__.py +++ /dev/null @@ -1,34 +0,0 @@ -"""File store module for persistent memory management. - -This module provides storage backends for memory chunks and file metadata, -including SQLite-based, ChromaDB-based, seekdb-based, and pure-Python local -implementations with vector and full-text search. -""" - -from .base_file_store import BaseFileStore -from .chroma_file_store import ChromaFileStore -from .local_file_store import LocalFileStore -from .sqlite_file_store import SqliteFileStore -from .zvec_file_store import ZvecFileStore -from ..registry_factory import R - -__all__ = [ - "BaseFileStore", - "ChromaFileStore", - "LocalFileStore", - "SqliteFileStore", - "ZvecFileStore", -] - -R.file_stores.register("sqlite")(SqliteFileStore) -R.file_stores.register("chroma")(ChromaFileStore) -R.file_stores.register("local")(LocalFileStore) -R.file_stores.register("zvec")(ZvecFileStore) - -try: - from .seekdb_file_store import SeekdbFileStore - - R.file_stores.register("seekdb")(SeekdbFileStore) - __all__.append("SeekdbFileStore") -except ImportError: - pass diff --git a/reme/core/file_store/base_file_store.py b/reme/core/file_store/base_file_store.py deleted file mode 100644 index 6a6b0901..00000000 --- a/reme/core/file_store/base_file_store.py +++ /dev/null @@ -1,227 +0,0 @@ -"""Base storage interface for file store.""" - -import re -from abc import ABC, abstractmethod -from pathlib import Path - -from ..embedding import BaseEmbeddingModel -from ..enumeration import MemorySource -from ..schema import FileMetadata, MemoryChunk, MemorySearchResult -from ..utils import get_logger - -logger = get_logger() - - -class BaseFileStore(ABC): - """Abstract base class for file storage backends.""" - - def __init__( - self, - store_name: str, - db_path: str | Path, - embedding_model: BaseEmbeddingModel | None = None, - vector_enabled: bool = False, - fts_enabled: bool = True, - **kwargs, - ): - """Initialize""" - # Validate store_name to prevent SQL injection - # Only allow alphanumeric characters and underscores - if not re.match(r"^[a-zA-Z0-9_]+$", store_name): - raise ValueError(f"Invalid '{store_name}'. Only alphanumeric characters and underscores are allowed.") - - # Ensure at least one search method is enabled - if not vector_enabled and not fts_enabled: - raise ValueError("At least one of vector_enabled or fts_enabled must be True.") - - # Ensure embedding_model is provided when vector search is enabled - if vector_enabled and embedding_model is None: - raise ValueError("embedding_model is required when vector_enabled is True.") - - self.store_name: str = store_name - self.db_path: Path = Path(db_path) - self.db_path.mkdir(parents=True, exist_ok=True) - self.embedding_model: BaseEmbeddingModel | None = embedding_model - self.vector_enabled: bool = vector_enabled - self.fts_enabled: bool = fts_enabled - self.kwargs: dict = kwargs - - @property - def embedding_dim(self) -> int: - """Get the embedding model's dimensionality.""" - if self.embedding_model is None: - return 1024 - return self.embedding_model.dimensions - - def _get_mock_embedding(self) -> list[float]: - """Generate a zero vector based on embedding model dimensions.""" - return [0.0] * self.embedding_dim - - def _disable_vector_search(self, reason: str = "embedding API error") -> None: - """Disable vector search and log a warning.""" - if self.vector_enabled: - logger.warning( - f"[{self.store_name}] Disabling vector search due to {reason}. " - "Falling back to full-text search only.", - ) - self.vector_enabled = False - - async def get_embedding(self, query: str, **kwargs) -> list[float]: - """Get embedding for a single query string.""" - if not self.vector_enabled: - return self._get_mock_embedding() - try: - return await self.embedding_model.get_embedding(query, **kwargs) - except Exception as e: - self._disable_vector_search(str(e)) - return self._get_mock_embedding() - - async def get_embeddings(self, queries: list[str], **kwargs) -> list[list[float]]: - """Get embeddings for a batch of query strings.""" - if not self.vector_enabled: - return [self._get_mock_embedding() for _ in queries] - try: - return await self.embedding_model.get_embeddings(queries, **kwargs) - except Exception as e: - self._disable_vector_search(str(e)) - return [self._get_mock_embedding() for _ in queries] - - async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk: - """Generate and populate embedding field for a single MemoryChunk object.""" - if not self.vector_enabled: - chunk.embedding = self._get_mock_embedding() - return chunk - try: - return await self.embedding_model.get_chunk_embedding(chunk, **kwargs) - except Exception as e: - self._disable_vector_search(str(e)) - chunk.embedding = self._get_mock_embedding() - return chunk - - async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]: - """Generate and populate embedding fields for a batch of MemoryChunk objects.""" - if not self.vector_enabled: - mock_embedding = self._get_mock_embedding() - for chunk in chunks: - chunk.embedding = mock_embedding.copy() - return chunks - try: - return await self.embedding_model.get_chunk_embeddings(chunks, **kwargs) - except Exception as e: - self._disable_vector_search(str(e)) - mock_embedding = self._get_mock_embedding() - for chunk in chunks: - chunk.embedding = mock_embedding.copy() - return chunks - - @abstractmethod - async def start(self): - """Initialize the storage backend.""" - - @abstractmethod - async def upsert_file(self, file_meta: FileMetadata, source: MemorySource, chunks: list[MemoryChunk]): - """Insert or update a file and its chunks.""" - - @abstractmethod - async def delete_file(self, path: str, source: MemorySource): - """Delete a file and all its chunks.""" - - @abstractmethod - async def delete_file_chunks(self, path: str, chunk_ids: list[str]): - """Delete chunks for a file.""" - - @abstractmethod - async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource): - """Insert or update specific chunks without affecting other chunks.""" - - @abstractmethod - async def list_files(self, source: MemorySource) -> list[str]: - """List all indexed file paths for a source.""" - - @abstractmethod - async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None: - """Get full file metadata with statistics.""" - - @abstractmethod - async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None: - """Update file metadata without affecting chunks. - - This is useful for incremental updates where only metadata needs to be updated - (e.g., after adding/removing chunks in delta file watcher). - - Args: - file_meta: Updated file metadata (hash, mtime_ms, size, chunk_count) - source: Memory source - """ - - @abstractmethod - async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]: - """Get all chunks for a file.""" - - @abstractmethod - async def vector_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform vector similarity search. - - Args: - query: Query embedding vector - limit: Maximum number of results - sources: Optional list of sources to filter - - Returns: - List of search results sorted by similarity - """ - - @abstractmethod - async def keyword_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform keyword/full-text search. - - Args: - query: Search query text - limit: Maximum number of results - sources: Optional list of sources to filter - - Returns: - List of search results sorted by relevance - """ - - @abstractmethod - async def hybrid_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - ) -> list[MemorySearchResult]: - """Perform hybrid search combining vector and keyword search. - - Args: - query: Search query text - limit: Maximum number of results - sources: Optional list of sources to filter - vector_weight: Weight for vector search results (0.0-1.0). - Keyword weight = 1.0 - vector_weight. - candidate_multiplier: Multiplier for candidate pool size. - candidates = limit * candidate_multiplier - - Returns: - List of search results sorted by combined relevance score - """ - - @abstractmethod - async def clear_all(self): - """Clear all indexed data.""" - - @abstractmethod - async def close(self): - """Close storage and release resources.""" diff --git a/reme/core/file_store/chroma_file_store.py b/reme/core/file_store/chroma_file_store.py deleted file mode 100644 index 6dda41b8..00000000 --- a/reme/core/file_store/chroma_file_store.py +++ /dev/null @@ -1,633 +0,0 @@ -"""ChromaDB storage backend for file store.""" - -import json -import random -import time -from pathlib import Path - -from .base_file_store import BaseFileStore -from ..enumeration import MemorySource -from ..schema import FileMetadata, MemoryChunk, MemorySearchResult -from ..utils import get_logger - -logger = get_logger() - -try: - import chromadb - from chromadb.config import Settings - - _CHROMADB_IMPORT_ERROR: Exception | None = None -except Exception as e: - _CHROMADB_IMPORT_ERROR = e - chromadb = None - Settings = None - - -class ChromaFileStore(BaseFileStore): - """ChromaDB file storage with vector and full-text search. - - Inherits embedding methods from BaseFileStore: - - get_chunk_embedding / get_chunk_embeddings (async) - - get_embedding / get_embeddings (async) - - Provides ChromaDB-backed persistent storage with: - - Vector similarity search (native ChromaDB) - - Full-text search (via ChromaDB where_document filter) - - Efficient chunk and file metadata management - """ - - def __init__( - self, - **kwargs, - ): - if _CHROMADB_IMPORT_ERROR is not None: - raise _CHROMADB_IMPORT_ERROR - - super().__init__(**kwargs) - self.client: "chromadb.ClientAPI | None" = None - self.chunks_collection: "chromadb.Collection | None" = None - # Initialize metadata file path (db_path and store_name are set by base class) - self._metadata_file: Path = self.db_path.parent / f"{self.store_name}_file_metadata.json" - self._metadata_cache: dict[str, dict[str, FileMetadata]] = {} - - @property - def collection_name(self) -> str: - """Get the name of the ChromaDB collection for this store.""" - return f"chunks_{self.store_name}" - - async def _load_metadata(self) -> dict[str, dict[str, FileMetadata]]: - """Load file metadata from disk. - - Returns: - Dictionary mapping source -> path -> FileMetadata - """ - if not self._metadata_file.exists(): - return {} - - try: - data = self._metadata_file.read_text(encoding="utf-8") - metadata_dict = json.loads(data) - - # Convert dict to FileMetadata objects - result = {} - for source, files in metadata_dict.items(): - result[source] = {} - for path, meta in files.items(): - result[source][path] = FileMetadata(**meta) - - logger.debug(f"Loaded file metadata from {self._metadata_file}") - return result - except Exception as e: - logger.warning(f"Failed to load file metadata from {self._metadata_file}: {e}") - return {} - - async def _save_metadata(self, metadata: dict[str, dict[str, FileMetadata]]) -> None: - """Save file metadata to disk. - - Args: - metadata: Dictionary mapping source -> path -> FileMetadata - """ - try: - # Convert FileMetadata objects to dict for JSON serialization - metadata_dict = {} - for source, files in metadata.items(): - metadata_dict[source] = {} - for path, meta in files.items(): - metadata_dict[source][path] = { - "path": meta.path, - "hash": meta.hash, - "mtime_ms": meta.mtime_ms, - "size": meta.size, - "chunk_count": meta.chunk_count, - } - - data = json.dumps(metadata_dict, indent=2, ensure_ascii=False) - self._metadata_file.write_text(data, encoding="utf-8") - logger.debug(f"Saved file metadata to {self._metadata_file}") - except Exception as e: - logger.error(f"Failed to save file metadata to {self._metadata_file}: {e}") - - async def start(self) -> None: - """Initialize ChromaDB client and collection.""" - if self.client is not None: - return - - # Initialize persistent ChromaDB client - self.client = chromadb.PersistentClient( - path=str(self.db_path), - settings=Settings( - anonymized_telemetry=False, - allow_reset=True, - ), - ) - - # Get or create the chunks collection - # ChromaDB uses cosine distance by default for similarity - self.chunks_collection = self.client.get_or_create_collection( - name=self.collection_name, - metadata={"hnsw:space": "cosine"}, - ) - - # Load metadata into cache - self._metadata_cache = await self._load_metadata() - - logger.info(f"ChromaDB initialized with collection: {self.collection_name}") - logger.info(f"File metadata will be persisted to: {self._metadata_file}") - - async def upsert_file( - self, - file_meta: FileMetadata, - source: MemorySource, - chunks: list[MemoryChunk], - ) -> None: - """Insert or update file and its chunks.""" - if not chunks: - return - - # Delete existing chunks for this file first - await self.delete_file(file_meta.path, source) - - # Batch generate embeddings for all chunks - # (base class returns mock embeddings when vector_enabled=False) - chunks = await self.get_chunk_embeddings(chunks) - - # Prepare data for ChromaDB batch upsert - ids = [] - documents = [] - embeddings = [] - metadatas = [] - - now = int(time.time() * 1000) - for chunk in chunks: - ids.append(chunk.id) - documents.append(chunk.text) - embeddings.append(chunk.embedding) - metadatas.append( - { - "path": file_meta.path, - "source": source.value, - "start_line": chunk.start_line, - "end_line": chunk.end_line, - "hash": chunk.hash, - "updated_at": now, - }, - ) - - # Batch upsert to ChromaDB (always pass embeddings to prevent default embedding function) - self.chunks_collection.upsert( - ids=ids, - documents=documents, - embeddings=embeddings, - metadatas=metadatas, - ) - - # Update file metadata in cache - if source.value not in self._metadata_cache: - self._metadata_cache[source.value] = {} - self._metadata_cache[source.value][file_meta.path] = FileMetadata( - hash=file_meta.hash, - mtime_ms=file_meta.mtime_ms, - size=file_meta.size, - path=file_meta.path, - chunk_count=len(chunks), - ) - - async def delete_file(self, path: str, source: MemorySource) -> None: - """Delete file and all its chunks.""" - # Query for all chunks with this path and source - results = self.chunks_collection.get( - where={"$and": [{"path": path}, {"source": source.value}]}, - include=[], - ) - - if results["ids"]: - self.chunks_collection.delete( - ids=results["ids"], - ) - - # Remove from file metadata cache - if source.value in self._metadata_cache: - self._metadata_cache[source.value].pop(path, None) - - async def delete_file_chunks(self, path: str, chunk_ids: list[str]) -> None: - """Delete specific chunks for a file.""" - if not chunk_ids: - return - - self.chunks_collection.delete( - ids=chunk_ids, - ) - - # Update chunk count in file metadata cache - for source_meta in self._metadata_cache.values(): - if path in source_meta: - # Recalculate chunk count - results = self.chunks_collection.get( - where={"path": path}, - include=[], - ) - source_meta[path].chunk_count = len(results["ids"]) - break - - async def upsert_chunks( - self, - chunks: list[MemoryChunk], - source: MemorySource, - ) -> None: - """Insert or update specific chunks without affecting other chunks.""" - if not chunks: - return - - # Batch generate embeddings for all chunks - # (base class returns mock embeddings when vector_enabled=False) - chunks = await self.get_chunk_embeddings(chunks) - - ids = [] - documents = [] - embeddings = [] - metadatas = [] - - now = int(time.time() * 1000) - for chunk in chunks: - ids.append(chunk.id) - documents.append(chunk.text) - embeddings.append(chunk.embedding) - metadatas.append( - { - "path": chunk.path, - "source": source.value, - "start_line": chunk.start_line, - "end_line": chunk.end_line, - "hash": chunk.hash, - "updated_at": now, - }, - ) - - # Always pass embeddings to prevent default embedding function - self.chunks_collection.upsert( - ids=ids, - documents=documents, - embeddings=embeddings, - metadatas=metadatas, - ) - - async def list_files(self, source: MemorySource) -> list[str]: - """List all indexed files for a source.""" - if source.value not in self._metadata_cache: - return [] - return list(self._metadata_cache[source.value].keys()) - - async def get_file_metadata( - self, - path: str, - source: MemorySource, - ) -> FileMetadata | None: - """Get file metadata with chunk count.""" - if source.value not in self._metadata_cache: - return None - return self._metadata_cache[source.value].get(path) - - async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None: - """Update file metadata without affecting chunks.""" - if source.value not in self._metadata_cache: - self._metadata_cache[source.value] = {} - - self._metadata_cache[source.value][file_meta.path] = FileMetadata( - hash=file_meta.hash, - mtime_ms=file_meta.mtime_ms, - size=file_meta.size, - path=file_meta.path, - chunk_count=file_meta.chunk_count, - ) - - async def get_file_chunks( - self, - path: str, - source: MemorySource, - ) -> list[MemoryChunk]: - """Get all chunks for a file.""" - results = self.chunks_collection.get( - where={"$and": [{"path": path}, {"source": source.value}]}, - include=["documents", "embeddings", "metadatas"], - ) - - chunks = [] - for i, chunk_id in enumerate(results["ids"]): - metadata = results["metadatas"][i] - chunks.append( - MemoryChunk( - id=chunk_id, - path=metadata["path"], - source=MemorySource(metadata["source"]), - start_line=metadata["start_line"], - end_line=metadata["end_line"], - text=results["documents"][i], - hash=metadata["hash"], - embedding=results["embeddings"][i] if results["embeddings"] is not None else None, - ), - ) - - # Sort by start_line - chunks.sort(key=lambda c: c.start_line) - return chunks - - async def vector_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform vector similarity search.""" - if not self.vector_enabled or not query: - return [] - - # Get query embedding - query_embedding = await self.get_embedding(query) - if not query_embedding: - return [] - - # Build where filter for sources - where_filter = None - if sources: - if len(sources) == 1: - where_filter = {"source": sources[0].value} - else: - where_filter = {"source": {"$in": [s.value for s in sources]}} - - # Perform vector search - try: - results = self.chunks_collection.query( - query_embeddings=[query_embedding], - n_results=limit, - where=where_filter, - include=["documents", "metadatas", "distances"], - ) - except Exception as e: - logger.error(f"Vector search failed: {e}, falling back to random results") - # Fallback: get some documents without vector search and assign random scores - try: - fallback_results = self.chunks_collection.get( - where=where_filter, - limit=limit, - include=["documents", "metadatas"], - ) - search_results = [] - if fallback_results["ids"]: - for i, _ in enumerate(fallback_results["ids"]): - metadata = fallback_results["metadatas"][i] - search_results.append( - MemorySearchResult( - path=metadata["path"], - start_line=metadata["start_line"], - end_line=metadata["end_line"], - score=random.uniform(0.3, 0.7), # Random score in middle range - snippet=fallback_results["documents"][i], - source=MemorySource(metadata["source"]), - raw_metric=None, - ), - ) - return search_results - except Exception as fallback_e: - logger.error(f"Fallback search also failed: {fallback_e}") - return [] - - search_results = [] - if results["ids"] and results["ids"][0]: - for i, _ in enumerate(results["ids"][0]): - metadata = results["metadatas"][0][i] - distance = results["distances"][0][i] - - # Convert cosine distance to similarity score - # Cosine distance range is [0, 2], convert to [1, 0] score - score = max(0.0, 1.0 - distance / 2.0) - - search_results.append( - MemorySearchResult( - path=metadata["path"], - start_line=metadata["start_line"], - end_line=metadata["end_line"], - score=score, - snippet=results["documents"][0][i], - source=MemorySource(metadata["source"]), - raw_metric=distance, - ), - ) - - # Sort by score descending - search_results.sort(key=lambda r: r.score, reverse=True) - return search_results - - async def keyword_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform keyword/full-text search. - - ChromaDB supports where_document filter for text matching. - Note: ChromaDB's $contains is case-sensitive, so we generate multiple - case variants (original, lowercase, capitalized) for each word to - improve recall while maintaining case-insensitive scoring. - """ - if not self.fts_enabled or not query: - return [] - - # Normalize whitespace and split into words - words = query.split() - if not words: - return [] - - # Generate case variants for each word to handle case-sensitive $contains - # Include: original, lowercase, and capitalized forms - word_variants = set() - for word in words: - word_variants.add(word) # original - word_variants.add(word.lower()) # lowercase - word_variants.add(word.capitalize()) # Capitalized - word_variants.add(word.upper()) # UPPERCASE - word_variants_list = list(word_variants) - - # Build where filter for sources - where_filter = None - if sources: - if len(sources) == 1: - where_filter = {"source": sources[0].value} - else: - where_filter = {"source": {"$in": [s.value for s in sources]}} - - # ChromaDB where_document uses $contains for substring matching (case-sensitive) - # Use multiple case variants to improve recall - if len(word_variants_list) == 1: - where_document: dict = {"$contains": word_variants_list[0]} - else: - where_document = {"$or": [{"$contains": w} for w in word_variants_list]} - - # Get all matching documents - results = self.chunks_collection.get( - where=where_filter, - where_document=where_document, - include=["documents", "metadatas"], - ) - - search_results = [] - query_lower = query.lower() - words_lower = [w.lower() for w in words] # lowercase words for scoring - n_words = len(words) - - for i, _ in enumerate(results["ids"]): - metadata = results["metadatas"][i] - text = results["documents"][i] - text_lower = text.lower() - - # Calculate relevance score based on word matches - match_count = sum(1 for w in words_lower if w in text_lower) - base_score = match_count / n_words - - # Bonus for full phrase match (only applies to multi-word queries) - phrase_bonus = 0.2 if n_words > 1 and query_lower in text_lower else 0.0 - # Scale base_score and add phrase bonus, max score is 1.0 - score = min(1.0, base_score + phrase_bonus) - - search_results.append( - MemorySearchResult( - path=metadata["path"], - start_line=metadata["start_line"], - end_line=metadata["end_line"], - score=score, - snippet=text, - source=MemorySource(metadata["source"]), - ), - ) - - # Sort by score descending and limit results - search_results.sort(key=lambda r: r.score, reverse=True) - return search_results[:limit] - - async def hybrid_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - ) -> list[MemorySearchResult]: - """Perform hybrid search combining vector and keyword search. - - Args: - query: Search query text - limit: Maximum number of results - sources: Optional list of sources to filter - vector_weight: Weight for vector search results (0.0-1.0). - Keyword weight = 1.0 - vector_weight. - candidate_multiplier: Multiplier for candidate pool size. - - Returns: - List of search results sorted by combined relevance score - """ - assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}" - - candidates = min(200, max(1, int(limit * candidate_multiplier))) - text_weight = 1.0 - vector_weight - - # Perform search based on enabled backends - if self.vector_enabled and self.fts_enabled: - keyword_results = await self.keyword_search(query, candidates, sources) - vector_results = await self.vector_search(query, candidates, sources) - - # Log original vector results - logger.info("\n=== Vector Search Results ===") - for i, r in enumerate(vector_results[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - # Log original keyword results - logger.info("\n=== Keyword Search Results ===") - for i, r in enumerate(keyword_results[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - if not keyword_results: - return vector_results[:limit] - elif not vector_results: - return keyword_results[:limit] - else: - merged = self._merge_hybrid_results( - vector=vector_results, - keyword=keyword_results, - vector_weight=vector_weight, - text_weight=text_weight, - ) - - # Log merged results - logger.info("\n=== Merged Hybrid Results ===") - for i, r in enumerate(merged[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - return merged[:limit] - elif self.vector_enabled: - vector_results = await self.vector_search(query, limit, sources) - return vector_results - elif self.fts_enabled: - keyword_results = await self.keyword_search(query, limit, sources) - return keyword_results - else: - return [] - - @staticmethod - def _merge_hybrid_results( - vector: list[MemorySearchResult], - keyword: list[MemorySearchResult], - vector_weight: float, - text_weight: float, - ) -> list[MemorySearchResult]: - """Merge vector and keyword search results with weighted scoring.""" - merged: dict[str, MemorySearchResult] = {} - - # Process vector results - for result in vector: - result.score = result.score * vector_weight - merged[result.merge_key] = result - - # Process keyword results - for result in keyword: - key = result.merge_key - if key in merged: - merged[key].score += result.score * text_weight - else: - result.score = result.score * text_weight - merged[key] = result - - # Sort by score and return - results = list(merged.values()) - results.sort(key=lambda r: r.score, reverse=True) - return results - - async def clear_all(self) -> None: - """Clear all indexed data.""" - # Delete and recreate the collection - self.client.delete_collection( - name=self.collection_name, - ) - self.chunks_collection = self.client.get_or_create_collection( - name=self.collection_name, - metadata={"hnsw:space": "cosine"}, - ) - - # Clear file metadata cache and disk - self._metadata_cache = {} - await self._save_metadata({}) - - logger.info(f"Cleared all data from ChromaDB collection: {self.collection_name}") - - async def close(self) -> None: - """Close ChromaDB client and release resources.""" - # Persist metadata cache to disk before closing - if self._metadata_cache: - await self._save_metadata(self._metadata_cache) - - # ChromaDB PersistentClient handles persistence automatically - self.client = None - self.chunks_collection = None - await super().close() diff --git a/reme/core/file_store/local_file_store.py b/reme/core/file_store/local_file_store.py deleted file mode 100644 index 09b73d8f..00000000 --- a/reme/core/file_store/local_file_store.py +++ /dev/null @@ -1,474 +0,0 @@ -"""Pure-Python in-memory storage backend for file store, with JSON file persistence.""" - -import json -from pathlib import Path - -import numpy as np -from loguru import logger - -from .base_file_store import BaseFileStore -from ..enumeration import MemorySource -from ..schema import FileMetadata, MemoryChunk, MemorySearchResult -from ..utils.common_utils import batch_cosine_similarity - - -class LocalFileStore(BaseFileStore): - """Pure-Python in-memory file storage with JSONL file persistence. - - No external dependencies required. All data lives in Python dicts; - writes are persisted to JSONL files on disk so state survives restarts. - - Inherits embedding methods from BaseFileStore: - - get_chunk_embedding / get_chunk_embeddings (async) - - get_embedding / get_embeddings (async) - - Provides: - - Vector similarity search (cosine similarity, pure Python) - - Full-text / keyword search (Python substring matching) - - Efficient chunk and file metadata management - """ - - def __init__(self, **kwargs): - super().__init__(**kwargs) - self._started: bool = False - # In-memory indexes - self._chunks: dict[str, MemoryChunk] = {} - self._files: dict[str, dict[str, FileMetadata]] = {} # source -> path -> meta - # Persistence paths (mirror ChromaFileStore convention) - self._chunks_file: Path = self.db_path / f"{self.store_name}_chunks.jsonl" - self._metadata_file: Path = self.db_path / f"{self.store_name}_file_metadata.json" - - # ------------------------------------------------------------------ - # Persistence helpers - # ------------------------------------------------------------------ - - async def _load_chunks(self) -> None: - """Load chunks from JSONL file into memory.""" - if not self._chunks_file.exists(): - return - try: - data = self._chunks_file.read_text(encoding="utf-8") - self._chunks = {} - for line in data.strip().split("\n"): - if not line: - continue - rec = json.loads(line) - chunk = MemoryChunk.model_validate(rec) - self._chunks[chunk.id] = chunk - logger.debug(f"Loaded {len(self._chunks)} chunks from {self._chunks_file}") - except Exception as e: - logger.warning(f"Failed to load chunks from {self._chunks_file}: {e}") - - async def _save_chunks(self) -> None: - """Persist chunks to JSONL file.""" - try: - lines = [] - for chunk in self._chunks.values(): - chunk_dict = chunk.model_dump(mode="json") - lines.append(json.dumps(chunk_dict, ensure_ascii=False)) - data = "\n".join(lines) - self._chunks_file.write_text(data, encoding="utf-8") - logger.debug(f"Saved {len(self._chunks)} chunks to {self._chunks_file}") - except Exception as e: - logger.error(f"Failed to save chunks to {self._chunks_file}: {e}") - - async def _load_metadata(self) -> None: - """Load file metadata from JSON file into memory.""" - if not self._metadata_file.exists(): - return - try: - data = self._metadata_file.read_text(encoding="utf-8") - raw: dict = json.loads(data) - self._files = { - source: {path: FileMetadata(**meta) for path, meta in files.items()} for source, files in raw.items() - } - logger.debug(f"Loaded file metadata from {self._metadata_file}") - except Exception as e: - logger.warning(f"Failed to load file metadata from {self._metadata_file}: {e}") - - async def _save_metadata(self) -> None: - """Persist file metadata to JSON file.""" - try: - raw: dict = {} - for source, files in self._files.items(): - raw[source] = { - path: { - "path": meta.path, - "hash": meta.hash, - "mtime_ms": meta.mtime_ms, - "size": meta.size, - "chunk_count": meta.chunk_count, - } - for path, meta in files.items() - } - data = json.dumps(raw, indent=2, ensure_ascii=False) - self._metadata_file.write_text(data, encoding="utf-8") - logger.debug(f"Saved file metadata to {self._metadata_file}") - except Exception as e: - logger.error(f"Failed to save file metadata to {self._metadata_file}: {e}") - - async def _persist(self) -> None: - """Persist all in-memory indexes after mutations.""" - await self._save_metadata() - await self._save_chunks() - - def _delete_file_in_memory(self, path: str, source: MemorySource) -> None: - """Delete file data from memory without flushing to disk.""" - to_delete = [cid for cid, chunk in self._chunks.items() if chunk.path == path and chunk.source == source] - for cid in to_delete: - del self._chunks[cid] - - if source.value in self._files: - self._files[source.value].pop(path, None) - - # ------------------------------------------------------------------ - # Lifecycle - # ------------------------------------------------------------------ - - async def start(self) -> None: - """Load persisted data into memory.""" - if self._started: - return - self._started = True - await self._load_metadata() - await self._load_chunks() - logger.info( - f"LocalFileStore '{self.store_name}' ready: " - f"{len(self._chunks)} chunks, metadata at {self._metadata_file}", - ) - - async def close(self) -> None: - """Flush state to disk and release memory.""" - await self._save_metadata() - await self._save_chunks() - self._chunks.clear() - self._files.clear() - self._started = False - - # ------------------------------------------------------------------ - # Write operations - # ------------------------------------------------------------------ - - async def upsert_file( - self, - file_meta: FileMetadata, - source: MemorySource, - chunks: list[MemoryChunk], - ) -> None: - """Insert or update file and its chunks.""" - if not chunks: - return - - # Remove existing chunks for this file/source first - self._delete_file_in_memory(file_meta.path, source) - - # Batch generate embeddings (base class returns mock embeddings when vector_enabled=False) - chunks = await self.get_chunk_embeddings(chunks) - - for chunk in chunks: - self._chunks[chunk.id] = chunk - - if source.value not in self._files: - self._files[source.value] = {} - self._files[source.value][file_meta.path] = FileMetadata( - hash=file_meta.hash, - mtime_ms=file_meta.mtime_ms, - size=file_meta.size, - path=file_meta.path, - chunk_count=len(chunks), - ) - await self._persist() - - async def delete_file(self, path: str, source: MemorySource) -> None: - """Delete file and all its chunks.""" - self._delete_file_in_memory(path, source) - await self._persist() - - async def delete_file_chunks(self, path: str, chunk_ids: list[str]) -> None: - """Delete specific chunks for a file.""" - if not chunk_ids: - return - - for cid in chunk_ids: - self._chunks.pop(cid, None) - - # Recalculate chunk_count in file metadata (per source) - for source_key, source_meta in self._files.items(): - if path in source_meta: - source_meta[path].chunk_count = sum( - 1 for chunk in self._chunks.values() if chunk.path == path and chunk.source.value == source_key - ) - await self._persist() - - async def upsert_chunks( - self, - chunks: list[MemoryChunk], - source: MemorySource, - ) -> None: - """Insert or update specific chunks without affecting other chunks.""" - if not chunks: - return - - chunks = await self.get_chunk_embeddings(chunks) - - for chunk in chunks: - self._chunks[chunk.id] = chunk - await self._persist() - - # ------------------------------------------------------------------ - # Read operations - # ------------------------------------------------------------------ - - async def list_files(self, source: MemorySource) -> list[str]: - """List all indexed files for a source.""" - return list(self._files.get(source.value, {}).keys()) - - async def get_file_metadata( - self, - path: str, - source: MemorySource, - ) -> FileMetadata | None: - """Get file metadata.""" - return self._files.get(source.value, {}).get(path) - - async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None: - """Update file metadata without affecting chunks.""" - if source.value not in self._files: - self._files[source.value] = {} - - self._files[source.value][file_meta.path] = FileMetadata( - hash=file_meta.hash, - mtime_ms=file_meta.mtime_ms, - size=file_meta.size, - path=file_meta.path, - chunk_count=file_meta.chunk_count, - ) - await self._persist() - - async def get_file_chunks( - self, - path: str, - source: MemorySource, - ) -> list[MemoryChunk]: - """Get all chunks for a file, sorted by start_line.""" - chunks = [chunk for chunk in self._chunks.values() if chunk.path == path and chunk.source == source] - chunks.sort(key=lambda c: c.start_line) - return chunks - - # ------------------------------------------------------------------ - # Search - # ------------------------------------------------------------------ - - async def vector_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform cosine-similarity vector search over in-memory embeddings.""" - if not self.vector_enabled or not query: - return [] - - query_embedding = await self.get_embedding(query) - if not query_embedding: - return [] - - expected_dim = self.embedding_dim - - # Collect candidate chunks with embeddings - candidates = [ - chunk for chunk in self._chunks.values() if (not sources or chunk.source in sources) and chunk.embedding - ] - - if not candidates: - return [] - - # Validate and fix chunk embedding dimensions - valid_embeddings = [] - for chunk in candidates: - emb = chunk.embedding - emb_len = len(emb) - if emb_len != expected_dim: - if emb_len < expected_dim: - emb = emb + [0.0] * (expected_dim - emb_len) - logger.warning( - f"Chunk embedding dimension {emb_len} < expected {expected_dim}, " - f"padded with zeros (chunk_id={chunk.id})", - ) - else: - emb = emb[:expected_dim] - logger.warning( - f"Chunk embedding dimension {emb_len} > expected {expected_dim}, " - f"truncated to {expected_dim} (chunk_id={chunk.id})", - ) - valid_embeddings.append(emb) - - # Build embedding matrix and compute similarities in batch - query_array = np.array([query_embedding]) # Shape: (1, emb_size) - chunk_embeddings = np.array(valid_embeddings) # Shape: (n, emb_size) - similarities = batch_cosine_similarity(query_array, chunk_embeddings)[0] # Shape: (n,) - - # Build results - results = [ - MemorySearchResult( - path=chunk.path, - start_line=chunk.start_line, - end_line=chunk.end_line, - score=float(similarity), - snippet=chunk.text, - source=chunk.source, - raw_metric=1.0 - float(similarity), - ) - for chunk, similarity in zip(candidates, similarities) - ] - - results.sort(key=lambda r: r.score, reverse=True) - return results[:limit] - - async def keyword_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform keyword/full-text search via Python substring matching.""" - if not self.fts_enabled or not query: - return [] - - words = query.split() - if not words: - return [] - - query_lower = query.lower() - words_lower = [w.lower() for w in words] - n_words = len(words) - - results = [] - for chunk in self._chunks.values(): - if sources and chunk.source not in sources: - continue - - text_lower = chunk.text.lower() - match_count = sum(1 for w in words_lower if w in text_lower) - if match_count == 0: - continue - - base_score = match_count / n_words - # Bonus for full phrase match (multi-word queries only) - phrase_bonus = 0.2 if n_words > 1 and query_lower in text_lower else 0.0 - score = min(1.0, base_score + phrase_bonus) - - results.append( - MemorySearchResult( - path=chunk.path, - start_line=chunk.start_line, - end_line=chunk.end_line, - score=score, - snippet=chunk.text, - source=chunk.source, - ), - ) - - results.sort(key=lambda r: r.score, reverse=True) - return results[:limit] - - async def hybrid_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - ) -> list[MemorySearchResult]: - """Perform hybrid search combining vector and keyword search. - - Args: - query: Search query text - limit: Maximum number of results - sources: Optional list of sources to filter - vector_weight: Weight for vector search results (0.0-1.0). - Keyword weight = 1.0 - vector_weight. - candidate_multiplier: Multiplier for candidate pool size. - - Returns: - List of search results sorted by combined relevance score - """ - assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}" - - candidates = min(200, max(1, int(limit * candidate_multiplier))) - text_weight = 1.0 - vector_weight - - if self.vector_enabled and self.fts_enabled: - keyword_results = await self.keyword_search(query, candidates, sources) - vector_results = await self.vector_search(query, candidates, sources) - - logger.info("\n=== Vector Search Results ===") - for i, r in enumerate(vector_results[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - logger.info("\n=== Keyword Search Results ===") - for i, r in enumerate(keyword_results[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - if not keyword_results: - return vector_results[:limit] - elif not vector_results: - return keyword_results[:limit] - else: - merged = self._merge_hybrid_results( - vector=vector_results, - keyword=keyword_results, - vector_weight=vector_weight, - text_weight=text_weight, - ) - - logger.info("\n=== Merged Hybrid Results ===") - for i, r in enumerate(merged[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - return merged[:limit] - elif self.vector_enabled: - return await self.vector_search(query, limit, sources) - elif self.fts_enabled: - return await self.keyword_search(query, limit, sources) - else: - return [] - - @staticmethod - def _merge_hybrid_results( - vector: list[MemorySearchResult], - keyword: list[MemorySearchResult], - vector_weight: float, - text_weight: float, - ) -> list[MemorySearchResult]: - """Merge vector and keyword search results with weighted scoring.""" - merged: dict[str, MemorySearchResult] = {} - - for result in vector: - result.metadata["_weighted_score"] = result.score * vector_weight - merged[result.merge_key] = result - - for result in keyword: - key = result.merge_key - if key in merged: - merged[key].metadata["_weighted_score"] += result.score * text_weight - else: - result.metadata["_weighted_score"] = result.score * text_weight - merged[key] = result - - results = list(merged.values()) - for r in results: - r.score = r.metadata.pop("_weighted_score") - - results.sort(key=lambda r: r.score, reverse=True) - return results - - async def clear_all(self) -> None: - """Clear all indexed data from memory and disk.""" - self._chunks.clear() - self._files.clear() - await self._persist() - logger.info(f"Cleared all data from LocalFileStore '{self.store_name}'") diff --git a/reme/core/file_store/seekdb_file_store.py b/reme/core/file_store/seekdb_file_store.py deleted file mode 100644 index 3f38400f..00000000 --- a/reme/core/file_store/seekdb_file_store.py +++ /dev/null @@ -1,601 +0,0 @@ -"""seekdb storage backend for file store. - -``pyseekdb.Client`` supports **embedded** (``path``) and **remote** OceanBase / -seekdb (``host`` / ``port`` / credentials). SQL-table-oriented helpers can still -use **pyobvector** via ``ObVecFileStore`` if needed. -""" - -import time -from pathlib import Path - -from loguru import logger - -from .base_file_store import BaseFileStore -from ..enumeration import MemorySource -from ..schema import FileMetadata, MemoryChunk, MemorySearchResult -from ..utils.pyseekdb_conn import ( - admin_kwargs_from_client_kwargs, - build_pyseekdb_client_kwargs, -) - -try: - import pyseekdb - from pyseekdb import Configuration, HNSWConfiguration, FulltextIndexConfig - - PYSEEKDB_AVAILABLE = True -except ImportError: - PYSEEKDB_AVAILABLE = False - pyseekdb = None - Configuration = None - HNSWConfiguration = None - FulltextIndexConfig = None - - -def _escape_sql(s: str) -> str: - """Escape single quotes for SQL string literals (MySQL/seekdb).""" - return s.replace("\\", "\\\\").replace("'", "''") - - -class SeekdbFileStore(BaseFileStore): - """seekdb file storage with vector and full-text search via ``pyseekdb``. - - **Embedded** (default): optional ``path`` for the data directory; if omitted, pyseekdb - uses its default (typically ``./seekdb.db``). **Remote**: ``host`` / ``port`` plus auth. - - File metadata is in a SQL table like SqliteFileStore; raw SQL uses ``execute``/``_execute``. - """ - - def __init__( - self, - host: str | None = None, - port: int | None = None, - user: str | None = None, - password: str = "", - path: str | None = None, - **kwargs, - ): - if not PYSEEKDB_AVAILABLE: - raise ImportError( - "pyseekdb is required for SeekdbFileStore. " - "Install it with: pip install reme-ai (pyseekdb is included)", - ) - - super().__init__(**kwargs) - self.client: "pyseekdb.Client | None" = None - self.collection = None - - self._is_remote, self._client_kw = build_pyseekdb_client_kwargs( - path=None if (host and host.strip()) else path, - database=self.store_name, - host=host, - port=port, - user=user, - password=password, - ) - - @property - def collection_name(self) -> str: - """Collection name for chunks.""" - return f"chunks_{self.store_name}" - - @property - def files_table_name(self) -> str: - """Table name for file metadata (same as SQLite).""" - return f"files_{self.store_name}" - - def _client_kwargs(self) -> dict: - """Kwargs for ``pyseekdb.Client`` (embedded or remote).""" - return self._client_kw - - def _sql_client(self): - """Underlying BaseClient for raw SQL (pyseekdb Client proxy exposes _server).""" - if self.client is None: - return None - return getattr(self.client, "_server", self.client) - - def _execute_sql(self, sql: str): - """Execute SQL via pyseekdb embedded/server client. Returns fetchall() for SELECT/SHOW/DESCRIBE.""" - client = self._sql_client() - if client is None: - raise RuntimeError("seekdb client not initialized") - for attr in ("execute", "_execute"): - run_sql = getattr(client, attr, None) - if run_sql is not None: - return run_sql(sql) - raise RuntimeError("seekdb client has no execute/_execute for raw SQL") - - def _create_files_table(self) -> None: - """Create file metadata table (same schema as SQLite).""" - sql = f""" - CREATE TABLE IF NOT EXISTS `{self.files_table_name}` ( - path VARCHAR(1024), - source VARCHAR(128), - hash VARCHAR(256), - mtime REAL, - size BIGINT, - PRIMARY KEY (path, source) - ) - """ - self._execute_sql(sql) - logger.debug(f"seekdb files table: {self.files_table_name}") - - async def start(self) -> None: - """Initialize seekdb client, collection, and files table.""" - if self.client is not None: - return - - kwargs = self._client_kwargs() - if not self._is_remote and "path" in kwargs: - Path(kwargs["path"]).parent.mkdir(parents=True, exist_ok=True) - database = kwargs.get("database", self.store_name) - try: - admin = pyseekdb.AdminClient(**admin_kwargs_from_client_kwargs(kwargs)) - if not any(db.name == database for db in admin.list_databases()): - admin.create_database(database) - except Exception as e: - logger.debug("seekdb AdminClient create_database: %s", e) - self.client = pyseekdb.Client(**kwargs) - - dim = self.embedding_dim - config = Configuration( - hnsw=HNSWConfiguration(dimension=dim, distance="cosine"), - fulltext_config=FulltextIndexConfig(analyzer="space"), - ) - self.collection = self.client.get_or_create_collection( - name=self.collection_name, - configuration=config, - embedding_function=None, - ) - self._create_files_table() - - logger.info( - f"seekdb initialized with collection: {self.collection_name}, " f"files table: {self.files_table_name}", - ) - - async def upsert_file( - self, - file_meta: FileMetadata, - source: MemorySource, - chunks: list[MemoryChunk], - ) -> None: - """Insert or update file and its chunks.""" - if not chunks: - return - - await self.delete_file(file_meta.path, source) - chunks = await self.get_chunk_embeddings(chunks) - - ids = [] - documents = [] - embeddings = [] - metadatas = [] - - now = int(time.time() * 1000) - for chunk in chunks: - ids.append(chunk.id) - documents.append(chunk.text) - embeddings.append(chunk.embedding) - metadatas.append( - { - "path": file_meta.path, - "source": source.value, - "start_line": chunk.start_line, - "end_line": chunk.end_line, - "hash": chunk.hash, - "updated_at": now, - }, - ) - - self.collection.upsert( - ids=ids, - documents=documents, - embeddings=embeddings, - metadatas=metadatas, - ) - # File metadata in DB table (same as SQLite) - p, s = _escape_sql(file_meta.path), _escape_sql(source.value) - h = _escape_sql(file_meta.hash) - mtime = file_meta.mtime_ms - size = file_meta.size - sql = ( - f"REPLACE INTO `{self.files_table_name}` (path, source, hash, mtime, size) " - f"VALUES ('{p}', '{s}', '{h}', {mtime}, {size})" - ) - self._execute_sql(sql) - - async def delete_file(self, path: str, source: MemorySource) -> None: - """Delete file and all its chunks.""" - results = self.collection.get( - where={"$and": [{"path": path}, {"source": source.value}]}, - include=[], - ) - if results.get("ids"): - self.collection.delete(ids=results["ids"]) - p, s = _escape_sql(path), _escape_sql(source.value) - self._execute_sql(f"DELETE FROM `{self.files_table_name}` WHERE path = '{p}' AND source = '{s}'") - - async def delete_file_chunks(self, path: str, chunk_ids: list[str]) -> None: - """Delete specific chunks for a file (chunk count comes from collection at get_file_metadata).""" - if not chunk_ids: - return - self.collection.delete(ids=chunk_ids) - - async def upsert_chunks( - self, - chunks: list[MemoryChunk], - source: MemorySource, - ) -> None: - """Insert or update specific chunks.""" - if not chunks: - return - chunks = await self.get_chunk_embeddings(chunks) - ids = [] - documents = [] - embeddings = [] - metadatas = [] - now = int(time.time() * 1000) - for chunk in chunks: - ids.append(chunk.id) - documents.append(chunk.text) - embeddings.append(chunk.embedding) - metadatas.append( - { - "path": chunk.path, - "source": source.value, - "start_line": chunk.start_line, - "end_line": chunk.end_line, - "hash": chunk.hash, - "updated_at": now, - }, - ) - self.collection.upsert( - ids=ids, - documents=documents, - embeddings=embeddings, - metadatas=metadatas, - ) - - async def list_files(self, source: MemorySource) -> list[str]: - """List all indexed files for a source (from files table).""" - s = _escape_sql(source.value) - rows = self._execute_sql(f"SELECT path FROM `{self.files_table_name}` WHERE source = '{s}'") - if not rows: - return [] - return [row[0] if isinstance(row, (list, tuple)) else row.get("path") for row in rows] - - async def get_file_metadata( - self, - path: str, - source: MemorySource, - ) -> FileMetadata | None: - """Get file metadata from files table and chunk count from collection.""" - p, s = _escape_sql(path), _escape_sql(source.value) - rows = self._execute_sql( - f"SELECT hash, mtime, size FROM `{self.files_table_name}` WHERE path = '{p}' AND source = '{s}'", - ) - if not rows: - return None - row = rows[0] - if isinstance(row, (list, tuple)): - hash_val, mtime, size = row[0], row[1], row[2] - else: - hash_val, mtime, size = row["hash"], row["mtime"], row["size"] - results = self.collection.get( - where={"$and": [{"path": path}, {"source": source.value}]}, - include=[], - ) - chunk_count = len(results.get("ids") or []) - return FileMetadata( - path=path, - hash=hash_val or "", - mtime_ms=mtime or 0, - size=size or 0, - chunk_count=chunk_count, - ) - - async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None: - """Update file metadata in files table without affecting chunks.""" - p = _escape_sql(file_meta.path) - s = _escape_sql(source.value) - h = _escape_sql(file_meta.hash) - mtime = file_meta.mtime_ms - size = file_meta.size - sql = ( - f"REPLACE INTO `{self.files_table_name}` (path, source, hash, mtime, size) " - f"VALUES ('{p}', '{s}', '{h}', {mtime}, {size})" - ) - self._execute_sql(sql) - - async def get_file_chunks( - self, - path: str, - source: MemorySource, - ) -> list[MemoryChunk]: - """Get all chunks for a file.""" - results = self.collection.get( - where={"$and": [{"path": path}, {"source": source.value}]}, - include=["documents", "embeddings", "metadatas"], - ) - chunks = [] - ids = results.get("ids") or [] - documents = results.get("documents") or [] - embeddings = results.get("embeddings") or [] - metadatas = results.get("metadatas") or [] - for i, chunk_id in enumerate(ids): - meta = metadatas[i] if i < len(metadatas) else {} - doc = documents[i] if i < len(documents) else "" - emb = embeddings[i] if embeddings and i < len(embeddings) else None - chunks.append( - MemoryChunk( - id=chunk_id, - path=meta.get("path", path), - source=MemorySource(meta.get("source", source.value)), - start_line=meta.get("start_line", 0), - end_line=meta.get("end_line", 0), - text=doc, - hash=meta.get("hash", ""), - embedding=emb, - ), - ) - chunks.sort(key=lambda c: c.start_line) - return chunks - - async def vector_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform vector similarity search.""" - if not self.vector_enabled or not query: - return [] - query_embedding = await self.get_embedding(query) - if not query_embedding: - return [] - - where_filter = None - if sources: - if len(sources) == 1: - where_filter = {"source": sources[0].value} - else: - where_filter = {"source": {"$in": [s.value for s in sources]}} - - results = self.collection.query( - query_embeddings=[query_embedding], - n_results=limit, - where=where_filter, - include=["documents", "metadatas", "distances"], - ) - - search_results = [] - if results.get("ids") and results["ids"][0]: - for i, _ in enumerate(results["ids"][0]): - metadata = results["metadatas"][0][i] - distance = results["distances"][0][i] - score = max(0.0, 1.0 - distance / 2.0) - search_results.append( - MemorySearchResult( - path=metadata["path"], - start_line=metadata["start_line"], - end_line=metadata["end_line"], - score=score, - snippet=results["documents"][0][i], - source=MemorySource(metadata["source"]), - raw_metric=distance, - ), - ) - search_results.sort(key=lambda r: r.score, reverse=True) - return search_results - - async def keyword_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform full-text keyword search.""" - if not self.fts_enabled or not query or not query.strip(): - return [] - - where_filter = None - if sources: - if len(sources) == 1: - where_filter = {"source": sources[0].value} - else: - where_filter = {"source": {"$in": [s.value for s in sources]}} - - results = self.collection.get( - where=where_filter, - where_document={"$contains": query.strip()}, - limit=limit, - include=["documents", "metadatas"], - ) - - search_results = [] - ids = results.get("ids") or [] - documents = results.get("documents") or [] - metadatas = results.get("metadatas") or [] - query_lower = query.lower() - words = query_lower.split() - n_words = max(1, len(words)) - - for i, _ in enumerate(ids): - meta = metadatas[i] if i < len(metadatas) else {} - text = documents[i] if i < len(documents) else "" - match_count = sum(1 for w in words if w in text.lower()) - score = min(1.0, match_count / n_words + (0.2 if query_lower in text.lower() else 0.0)) - search_results.append( - MemorySearchResult( - path=meta.get("path", ""), - start_line=meta.get("start_line", 0), - end_line=meta.get("end_line", 0), - score=score, - snippet=text, - source=MemorySource(meta.get("source", "memory")), - ), - ) - search_results.sort(key=lambda r: r.score, reverse=True) - return search_results[:limit] - - async def hybrid_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - ) -> list[MemorySearchResult]: - """Perform hybrid search using seekdb native hybrid_search when possible.""" - if not query or not query.strip(): - return [] - - assert 0.0 <= vector_weight <= 1.0 - candidates = min(200, max(1, int(limit * candidate_multiplier))) - where_filter = None - if sources: - if len(sources) == 1: - where_filter = {"source": sources[0].value} - else: - where_filter = {"source": {"$in": [s.value for s in sources]}} - - if self.vector_enabled and self.fts_enabled: - try: - query_embedding = await self.get_embedding(query) - if not query_embedding: - return await self.keyword_search(query, limit, sources) - - results = self.collection.hybrid_search( - query={ - "where_document": {"$contains": query.strip()}, - "where": where_filter, - "n_results": candidates, - }, - knn={ - "query_embeddings": [query_embedding], - "where": where_filter, - "n_results": candidates, - }, - rank={"rrf": {}}, - n_results=limit, - include=["documents", "metadatas", "distances"], - ) - except Exception as e: - logger.warning(f"seekdb hybrid_search failed, fallback to merge: {e}") - return await self._hybrid_search_merge( - query, - limit, - sources, - vector_weight, - candidate_multiplier, - ) - else: - return await self._hybrid_search_merge( - query, - limit, - sources, - vector_weight, - candidate_multiplier, - ) - - search_results = [] - raw_ids = results.get("ids") or [] - ids = raw_ids[0] if raw_ids and isinstance(raw_ids[0], list) else raw_ids - if not ids: - return [] - documents = results.get("documents") - metadatas = results.get("metadatas") - distances = results.get("distances") - doc_list = (documents[0] if documents and isinstance(documents[0], list) else documents) or [] - meta_list = (metadatas[0] if metadatas and isinstance(metadatas[0], list) else metadatas) or [] - dist_list = (distances[0] if distances and isinstance(distances[0], list) else distances) or [] - for i, _ in enumerate(ids): - meta = meta_list[i] if i < len(meta_list) else {} - doc = doc_list[i] if i < len(doc_list) else "" - dist = dist_list[i] if i < len(dist_list) else 0.0 - score = max(0.0, 1.0 - dist / 2.0) - search_results.append( - MemorySearchResult( - path=meta.get("path", ""), - start_line=meta.get("start_line", 0), - end_line=meta.get("end_line", 0), - score=score, - snippet=doc, - source=MemorySource(meta.get("source", "memory")), - raw_metric=dist, - ), - ) - return search_results[:limit] - - async def _hybrid_search_merge( - self, - query: str, - limit: int, - sources: list[MemorySource] | None, - vector_weight: float, - candidate_multiplier: float, - ) -> list[MemorySearchResult]: - """Fallback: merge vector and keyword results like Chroma.""" - candidates = min(200, max(1, int(limit * candidate_multiplier))) - text_weight = 1.0 - vector_weight - - if self.vector_enabled and self.fts_enabled: - keyword_results = await self.keyword_search(query, candidates, sources) - vector_results = await self.vector_search(query, candidates, sources) - if not keyword_results: - return vector_results[:limit] - if not vector_results: - return keyword_results[:limit] - merged = self._merge_hybrid_results( - vector=vector_results, - keyword=keyword_results, - vector_weight=vector_weight, - text_weight=text_weight, - ) - return merged[:limit] - if self.vector_enabled: - return await self.vector_search(query, limit, sources) - return await self.keyword_search(query, limit, sources) - - @staticmethod - def _merge_hybrid_results( - vector: list[MemorySearchResult], - keyword: list[MemorySearchResult], - vector_weight: float, - text_weight: float, - ) -> list[MemorySearchResult]: - """Merge vector and keyword results with weighted scoring.""" - merged: dict[str, MemorySearchResult] = {} - for result in vector: - result.score = result.score * vector_weight - merged[result.merge_key] = result - for result in keyword: - key = result.merge_key - if key in merged: - merged[key].score += result.score * text_weight - else: - result.score = result.score * text_weight - merged[key] = result - results = list(merged.values()) - results.sort(key=lambda r: r.score, reverse=True) - return results - - async def clear_all(self) -> None: - """Clear all indexed data (collection + files table).""" - self.client.delete_collection(self.collection_name) - dim = self.embedding_dim - config = Configuration( - hnsw=HNSWConfiguration(dimension=dim, distance="cosine"), - fulltext_config=FulltextIndexConfig(analyzer="space"), - ) - self.collection = self.client.get_or_create_collection( - name=self.collection_name, - configuration=config, - embedding_function=None, - ) - self._execute_sql(f"DELETE FROM `{self.files_table_name}`") - logger.info(f"Cleared all data from seekdb collection: {self.collection_name} and files table") - - async def close(self) -> None: - """Close client (file metadata is in DB, no persist needed).""" - self.client = None - self.collection = None diff --git a/reme/core/file_store/sqlite_file_store.py b/reme/core/file_store/sqlite_file_store.py deleted file mode 100644 index 0a494d2c..00000000 --- a/reme/core/file_store/sqlite_file_store.py +++ /dev/null @@ -1,978 +0,0 @@ -"""SQLite storage backend for file store.""" - -import json - -import struct -import time - -from loguru import logger - -from .base_file_store import BaseFileStore -from ..enumeration import MemorySource -from ..schema import FileMetadata, MemoryChunk, MemorySearchResult - - -class SqliteFileStore(BaseFileStore): - """SQLite file storage with vector and full-text search. - - Inherits embedding methods from BaseFileStore: - - get_chunk_embedding / get_chunk_embeddings (async) - - get_chunk_embedding_sync / get_chunk_embeddings_sync (sync) - - get_embedding / get_embeddings (async) - - Provides SQLite-backed persistent storage with: - - Vector similarity search (via sqlite-vec extension) - - Full-text search (via FTS5) - - Efficient chunk and file metadata management - """ - - def __init__(self, vec_ext_path: str = "", **kwargs): - super().__init__(**kwargs) - self.vec_ext_path = vec_ext_path - import sqlite3 - - self.conn: sqlite3.Connection | None = None - - @property - def vector_table_name(self) -> str: - """Get the name of the vector table for this store.""" - return f"chunks_vec_{self.store_name}" - - @property - def fts_table_name(self) -> str: - """Get the name of the FTS table for this store.""" - return f"chunks_fts_{self.store_name}" - - @property - def chunks_table_name(self) -> str: - """Get the name of the chunks table for this store.""" - return f"chunks_{self.store_name}" - - @property - def files_table_name(self) -> str: - """Get the name of the files table for this store.""" - return f"files_{self.store_name}" - - @staticmethod - def vector_to_blob(embedding: list[float]) -> bytes: - """Convert vector to binary blob for sqlite-vec.""" - return struct.pack(f"{len(embedding)}f", *embedding) - - async def start(self) -> None: - """Initialize database and load extensions.""" - if self.conn is not None: - return - import sqlite3 - - self.conn = sqlite3.connect(self.db_path / "reme.db", check_same_thread=False) - - # Only load sqlite-vec extension if vector search is enabled - if self.vector_enabled: - logger.warning( - "On macOS systems with version 14 or earlier, " - "loading the sqlite-vec vector extension carries a risk of crashes or hangs.", - ) - - self.conn.enable_load_extension(True) - - # Load sqlite-vec extension - if self.vec_ext_path: - try: - self.conn.load_extension(self.vec_ext_path) - logger.info(f"Loaded sqlite-vec: {self.vec_ext_path}") - except Exception as e: - logger.warning(f"Failed to load sqlite-vec: {e}") - - else: - try: - import sqlite_vec - - ext_path = sqlite_vec.loadable_path() - self.conn.load_extension(ext_path) - logger.info(f"Loaded sqlite-vec from package: {ext_path}") - - except Exception as e: - logger.warning(f"Failed to load sqlite-vec from package: {e}") - # Fallback: try common extension names - for name in ["vec0", "sqlite_vec", "vector0"]: - try: - self.conn.load_extension(name) - logger.info(f"Loaded sqlite-vec: {name}") - break - except Exception: - pass - - self.conn.enable_load_extension(False) - else: - logger.info("Vector search disabled, skipping sqlite-vec extension loading") - - await self._create_tables() - - async def _create_tables(self) -> None: - """Create database schema.""" - cursor = self.conn.cursor() - try: - # Files - cursor.execute( - f""" - CREATE TABLE IF NOT EXISTS {self.files_table_name} ( - path TEXT, - source TEXT, - hash TEXT, - mtime REAL, - size INTEGER, - PRIMARY KEY (path, source) - ) - """, - ) - - # Chunks - cursor.execute( - f""" - CREATE TABLE IF NOT EXISTS {self.chunks_table_name} ( - id TEXT PRIMARY KEY, - path TEXT, - source TEXT, - start_line INTEGER, - end_line INTEGER, - hash TEXT, - text TEXT, - embedding TEXT, - updated_at INTEGER - ) - """, - ) - - # Vector table (sqlite-vec) - if self.vector_enabled: - cursor.execute( - f""" - CREATE VIRTUAL TABLE IF NOT EXISTS {self.vector_table_name} USING vec0( - id TEXT PRIMARY KEY, - embedding FLOAT[{self.embedding_dim}] - ) - """, - ) - logger.info(f"Created vector table (dims={self.embedding_dim})") - - # FTS table - if self.fts_enabled: - cursor.execute( - f""" - CREATE VIRTUAL TABLE IF NOT EXISTS {self.fts_table_name} USING fts5( - text, - id UNINDEXED, - path UNINDEXED, - source UNINDEXED, - start_line UNINDEXED, - end_line UNINDEXED, - tokenize='trigram' - ) - """, - ) - logger.info("Created FTS5 table with trigram tokenizer") - - self.conn.commit() - except Exception as e: - logger.error(f"Failed to create tables: {e}") - raise - finally: - cursor.close() - - async def upsert_file(self, file_meta: FileMetadata, source: MemorySource, chunks: list[MemoryChunk]): - """Insert or update file and its chunks.""" - cursor = self.conn.cursor() - - try: - cursor.execute("BEGIN") - - # Insert file - cursor.execute( - f""" - INSERT OR REPLACE INTO {self.files_table_name} (path, source, hash, mtime, size) - VALUES (?, ?, ?, ?, ?) - """, - (file_meta.path, source.value, file_meta.hash, file_meta.mtime_ms, file_meta.size), - ) - - # Insert chunks - now = int(time.time() * 1000) - for chunk in chunks: - cursor.execute( - f""" - INSERT OR REPLACE INTO {self.chunks_table_name} ( - id, path, source, start_line, end_line, - hash, text, embedding, updated_at - ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) - """, - ( - chunk.id, - file_meta.path, - source.value, - chunk.start_line, - chunk.end_line, - chunk.hash, - chunk.text, - json.dumps(chunk.embedding) if chunk.embedding else None, - now, - ), - ) - - # Insert vector (vec0 doesn't support OR REPLACE, use DELETE + INSERT) - if self.vector_enabled: - if not chunk.embedding: - logger.warning( - f"Chunk {chunk.id} missing embedding for vector insert, skipping vector indexing", - ) - else: - # Delete existing vector first - cursor.execute( - f"DELETE FROM {self.vector_table_name} WHERE id = ?", - (chunk.id,), - ) - # Then insert new vector - cursor.execute( - f"INSERT INTO {self.vector_table_name} (id, embedding) VALUES (?, ?)", - (chunk.id, self.vector_to_blob(chunk.embedding)), - ) - - # Insert FTS - if self.fts_enabled: - cursor.execute( - f""" - INSERT OR REPLACE INTO {self.fts_table_name} ( - text, id, path, source, start_line, end_line - ) VALUES (?, ?, ?, ?, ?, ?) - """, - ( - chunk.text, - chunk.id, - file_meta.path, - source.value, - chunk.start_line, - chunk.end_line, - ), - ) - - cursor.execute("COMMIT") - except Exception as e: - cursor.execute("ROLLBACK") - logger.error(f"Failed to upsert file {file_meta.path}: {e}") - raise - finally: - cursor.close() - - async def delete_file(self, path: str, source: MemorySource): - """Delete file and all its chunks.""" - cursor = self.conn.cursor() - try: - cursor.execute("BEGIN") - - # Get chunk IDs for vector deletion - cursor.execute( - f"SELECT id FROM {self.chunks_table_name} WHERE path = ? AND source = ?", - (path, source.value), - ) - chunk_ids = [row[0] for row in cursor.fetchall()] - - # Delete vectors - if self.vector_enabled and chunk_ids: - for chunk_id in chunk_ids: - cursor.execute( - f"DELETE FROM {self.vector_table_name} WHERE id = ?", - (chunk_id,), - ) - - # Delete FTS entries - if self.fts_enabled: - cursor.execute( - f"DELETE FROM {self.fts_table_name} WHERE path = ? AND source = ?", - (path, source.value), - ) - - # Delete chunks and file - cursor.execute( - f"DELETE FROM {self.chunks_table_name} WHERE path = ? AND source = ?", - (path, source.value), - ) - cursor.execute( - f"DELETE FROM {self.files_table_name} WHERE path = ? AND source = ?", - (path, source.value), - ) - - cursor.execute("COMMIT") - except Exception as e: - cursor.execute("ROLLBACK") - logger.error(f"Failed to delete file {path}: {e}") - raise - finally: - cursor.close() - - async def delete_file_chunks(self, path: str, chunk_ids: list[str]): - """Delete specific chunks for a file.""" - if not chunk_ids: - return - - cursor = self.conn.cursor() - try: - cursor.execute("BEGIN") - - # Delete vectors - if self.vector_enabled: - for chunk_id in chunk_ids: - cursor.execute( - f"DELETE FROM {self.vector_table_name} WHERE id = ?", - (chunk_id,), - ) - - # Delete FTS entries - if self.fts_enabled: - placeholders = ",".join("?" * len(chunk_ids)) - cursor.execute( - f"DELETE FROM {self.fts_table_name} WHERE id IN ({placeholders})", - chunk_ids, - ) - - # Delete chunks - placeholders = ",".join("?" * len(chunk_ids)) - cursor.execute( - f"DELETE FROM {self.chunks_table_name} WHERE id IN ({placeholders})", - chunk_ids, - ) - - cursor.execute("COMMIT") - except Exception as e: - cursor.execute("ROLLBACK") - logger.error(f"Failed to delete chunks for {path}: {e}") - raise - finally: - cursor.close() - - async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource): - """Insert or update specific chunks without affecting other chunks.""" - if not chunks: - return - - cursor = self.conn.cursor() - try: - cursor.execute("BEGIN") - - now = int(time.time() * 1000) - for chunk in chunks: - # Insert/update chunk - cursor.execute( - f""" - INSERT OR REPLACE INTO {self.chunks_table_name} ( - id, path, source, start_line, end_line, - hash, text, embedding, updated_at - ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) - """, - ( - chunk.id, - chunk.path, - source.value, - chunk.start_line, - chunk.end_line, - chunk.hash, - chunk.text, - json.dumps(chunk.embedding) if chunk.embedding else None, - now, - ), - ) - - # Insert/update vector (vec0 doesn't support OR REPLACE, use DELETE + INSERT) - if self.vector_enabled: - if not chunk.embedding: - logger.warning( - f"Chunk {chunk.id} missing embedding for vector insert, skipping vector indexing", - ) - else: - # Delete existing vector first - cursor.execute( - f"DELETE FROM {self.vector_table_name} WHERE id = ?", - (chunk.id,), - ) - # Then insert new vector - cursor.execute( - f"INSERT INTO {self.vector_table_name} (id, embedding) VALUES (?, ?)", - (chunk.id, self.vector_to_blob(chunk.embedding)), - ) - - # Insert/update FTS - if self.fts_enabled: - cursor.execute( - f""" - INSERT OR REPLACE INTO {self.fts_table_name} ( - text, id, path, source, start_line, end_line - ) VALUES (?, ?, ?, ?, ?, ?) - """, - ( - chunk.text, - chunk.id, - chunk.path, - source.value, - chunk.start_line, - chunk.end_line, - ), - ) - - cursor.execute("COMMIT") - except Exception as e: - cursor.execute("ROLLBACK") - logger.error(f"Failed to upsert chunks: {e}") - raise - finally: - cursor.close() - - async def list_files(self, source: MemorySource) -> list[str]: - """List all indexed files.""" - cursor = self.conn.cursor() - try: - cursor.execute(f"SELECT path FROM {self.files_table_name} WHERE source = ?", (source.value,)) - paths = [row[0] for row in cursor.fetchall()] - return paths - except Exception as e: - logger.error(f"Failed to list files: {e}") - raise - finally: - cursor.close() - - async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None: - """Get file metadata with chunk count.""" - cursor = self.conn.cursor() - try: - cursor.execute( - f"SELECT hash, mtime, size FROM {self.files_table_name} WHERE path = ? AND source = ?", - (path, source.value), - ) - row = cursor.fetchone() - if not row: - return None - - hash_val, mtime, size = row - cursor.execute( - f"SELECT COUNT(*) FROM {self.chunks_table_name} WHERE path = ? AND source = ?", - (path, source.value), - ) - chunk_count = cursor.fetchone()[0] - - return FileMetadata( - hash=hash_val, - mtime_ms=mtime, - size=size, - path=path, - chunk_count=chunk_count, - ) - except Exception as e: - logger.error(f"Failed to get file metadata for {path}: {e}") - raise - finally: - cursor.close() - - async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None: - """Update file metadata without affecting chunks.""" - cursor = self.conn.cursor() - try: - cursor.execute( - f""" - INSERT OR REPLACE INTO {self.files_table_name} (path, source, hash, mtime, size) - VALUES (?, ?, ?, ?, ?) - """, - (file_meta.path, source.value, file_meta.hash, file_meta.mtime_ms, file_meta.size), - ) - self.conn.commit() - except Exception as e: - logger.error(f"Failed to update file metadata for {file_meta.path}: {e}") - raise - finally: - cursor.close() - - async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]: - """Get all chunks for a file.""" - cursor = self.conn.cursor() - try: - cursor.execute( - f""" - SELECT id, path, source, start_line, end_line, text, hash, embedding - FROM {self.chunks_table_name} WHERE path = ? AND source = ? - ORDER BY start_line - """, - (path, source.value), - ) - - chunks = [] - for row in cursor.fetchall(): - chunk_id, path_val, source_val, start, end, text, hash_val, emb_str = row - # Parse embedding from JSON string - embedding = None - if emb_str: - try: - embedding = json.loads(emb_str) - except (json.JSONDecodeError, TypeError): - embedding = None - - chunks.append( - MemoryChunk( - id=chunk_id, - path=path_val, - source=MemorySource(source_val), - start_line=start, - end_line=end, - text=text, - hash=hash_val, - embedding=embedding, - ), - ) - - return chunks - except Exception as e: - logger.error(f"Failed to get file chunks for {path}: {e}") - raise - finally: - cursor.close() - - async def vector_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform vector similarity search.""" - if not self.vector_enabled or not query: - return [] - - # Get query embedding - query_embedding = await self.get_embedding(query) - if not query_embedding: - return [] - - cursor = self.conn.cursor() - source_filter = "" - params: list = [] - if sources: - placeholders = ",".join("?" * len(sources)) - source_filter = f" AND c.source IN ({placeholders})" - params = [s.value for s in sources] - - try: - query_blob = self.vector_to_blob(query_embedding) - - # Correct SQLite-vec syntax for vector search with limit - # vec0 requires 'k = ?' constraint for knn queries - query_sql = f""" - SELECT c.id, c.path, c.start_line, c.end_line, c.source, c.text, v.distance - FROM {self.vector_table_name} v - JOIN {self.chunks_table_name} c ON v.id = c.id - WHERE v.embedding MATCH ? - AND k = ? - """ - query_params: list = [query_blob, limit] - - # Add source filter if specified - if source_filter: - query_sql += source_filter - query_params.extend(params) - - # Order by distance (k constraint already limits results) - query_sql += " ORDER BY v.distance" - - cursor.execute(query_sql, query_params) - - results = [] - for _, path, start, end, src, text, dist in cursor.fetchall(): - # Convert L2 distance to similarity score - # For normalized vectors, L2 distance range is [0, 2] - # Map to [1, 0] score range (higher score = more similar) - score = max(0.0, 1.0 - dist / 2.0) - snippet = text - results.append( - MemorySearchResult( - path=path, - start_line=start, - end_line=end, - score=score, - snippet=snippet, - source=MemorySource(src), - raw_metric=dist, - ), - ) - - results.sort(key=lambda r: r.score, reverse=True) - return results - except Exception as e: - logger.error(f"Vector search failed: {e}") - return [] - finally: - cursor.close() - - @staticmethod - def _sanitize_fts_query(query: str) -> str: - """Sanitize query string for FTS5 search. - - Removes or escapes special characters that have special meaning in FTS5: - - * (prefix match) - - ? (not used in FTS5, but can cause issues) - - " (phrase search, needs escaping) - - : (column filter) - - ^ (start of line anchor, not standard FTS5) - - ' (single quote, causes syntax errors) - - ` (backtick, can cause issues) - - | (pipe, OR operator) - - + (plus, can be used for required terms) - - - (minus, NOT operator) - - = (equals, can cause issues) - - < > (angle brackets, comparison operators) - - ! (exclamation, NOT operator variant) - - @ # $ % & (other special chars) - - "\" - - / (slash, can interfere) - - ; (semicolon, statement separator) - - , (comma, can interfere with phrase parsing) - - Args: - query: Raw query string - - Returns: - Sanitized query string safe for FTS5 - """ - if not query: - return "" - - # Remove FTS5 special characters that we don't want users to use - # Keep only alphanumeric, spaces, periods, and underscores - special_chars = [ - "*", - "?", - ":", - "^", - "(", - ")", - "[", - "]", - "{", - "}", - "'", - '"', - "`", - "|", - "+", - "-", - "=", - "<", - ">", - "!", - "@", - "#", - "$", - "%", - "&", - "\\", - "/", - ";", - ",", - ] - cleaned = query - for char in special_chars: - cleaned = cleaned.replace(char, " ") - - # Normalize whitespace - cleaned = " ".join(cleaned.split()) - - return cleaned - - async def keyword_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform keyword search. - - Strategy: - - FTS5 trigram (fast path): used when ALL terms >= 3 chars (trigram minimum). - - LIKE (universal fallback): used when any term < 3 chars, covering CJK - short words, single/double-char queries, and mixed-length queries. - """ - if not self.fts_enabled: - return [] - - cleaned = self._sanitize_fts_query(query) - if not cleaned: - return [] - - words = cleaned.split() - if not words: - return [] - - # FTS5 trigram requires every term >= 3 characters - if all(len(w) >= 3 for w in words): - results = await self._fts_trigram_search(words, limit, sources) - if results: - return results - - # Universal fallback: LIKE-based substring search - return await self._like_search(cleaned, words, limit, sources) - - async def _fts_trigram_search( - self, - words: list[str], - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """FTS5 trigram search. All terms must be >= 3 characters.""" - escaped_words = [w.replace('"', '""') for w in words] - fts_query = " OR ".join(escaped_words) - - cursor = self.conn.cursor() - source_filter = "" - params: list = [fts_query] - if sources: - placeholders = ",".join("?" * len(sources)) - source_filter = f" AND fts.source IN ({placeholders})" - params.extend([s.value for s in sources]) - params.append(limit) - - try: - cursor.execute( - f""" - SELECT fts.id, fts.path, fts.start_line, fts.end_line, - fts.source, fts.text, rank - FROM {self.fts_table_name} fts - WHERE fts.text MATCH ?{source_filter} - ORDER BY rank - LIMIT ? - """, - params, - ) - - results = [] - for _, path, start, end, src, text, rank in cursor.fetchall(): - score = max(0.0, 1.0 / (1.0 + abs(rank))) - results.append( - MemorySearchResult( - path=path, - start_line=start, - end_line=end, - score=score, - snippet=text, - source=MemorySource(src), - raw_metric=rank, - ), - ) - results.sort(key=lambda r: r.score, reverse=True) - return results - except Exception as e: - logger.error(f"FTS trigram search failed: {e}") - return [] - finally: - cursor.close() - - async def _like_search( - self, - phrase: str, - words: list[str], - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """LIKE-based substring search with Python-side relevance scoring. - - Handles any term length and all languages (CJK, Latin, etc.). - Scores results by: word-match ratio + full-phrase bonus. - """ - cursor = self.conn.cursor() - try: - # Build OR conditions: match any individual word - like_clauses = [] - params: list = [] - for word in words: - like_clauses.append("c.text LIKE ?") - params.append(f"%{word}%") - - where_clause = " OR ".join(like_clauses) - - source_filter = "" - if sources: - placeholders = ",".join("?" * len(sources)) - source_filter = f" AND c.source IN ({placeholders})" - params.extend([s.value for s in sources]) - - # Fetch extra candidates for re-ranking in Python - fetch_limit = min(limit * 3, 200) - params.append(fetch_limit) - - cursor.execute( - f""" - SELECT c.id, c.path, c.start_line, c.end_line, c.source, c.text - FROM {self.chunks_table_name} c - WHERE ({where_clause}){source_filter} - LIMIT ? - """, - params, - ) - - results = [] - phrase_lower = phrase.lower() - words_lower = [w.lower() for w in words] - n_words = len(words) - - for _, path, start, end, src, text in cursor.fetchall(): - text_lower = text.lower() - - # Base score: proportion of query words found in text - match_count = sum(1 for w in words_lower if w in text_lower) - base_score = match_count / n_words - - # Bonus: full phrase appears as contiguous substring - phrase_bonus = 0.2 if n_words > 1 and phrase_lower in text_lower else 0.0 - - score = min(1.0, base_score * 0.8 + phrase_bonus) - - results.append( - MemorySearchResult( - path=path, - start_line=start, - end_line=end, - score=score, - snippet=text, - source=MemorySource(src), - ), - ) - - # Sort by score descending, return top `limit` - results.sort(key=lambda r: r.score, reverse=True) - return results[:limit] - except Exception as e: - logger.error(f"LIKE search failed: {e}") - return [] - finally: - cursor.close() - - async def hybrid_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - ) -> list[MemorySearchResult]: - """Perform hybrid search combining vector and keyword search. - - Args: - query: Search query text - limit: Maximum number of results - sources: Optional list of sources to filter - vector_weight: Weight for vector search results (0.0-1.0). - Keyword weight = 1.0 - vector_weight. - candidate_multiplier: Multiplier for candidate pool size. - - Returns: - List of search results sorted by combined relevance score - """ - assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}" - - candidates = min(200, max(1, int(limit * candidate_multiplier))) - text_weight = 1.0 - vector_weight - - # Perform search based on enabled backends - if self.vector_enabled and self.fts_enabled: - keyword_results = await self.keyword_search(query, candidates, sources) - vector_results = await self.vector_search(query, candidates, sources) - - # Log original vector results - logger.info("\n=== Vector Search Results ===") - for i, r in enumerate(vector_results[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - # Log original keyword results - logger.info("\n=== Keyword Search Results ===") - for i, r in enumerate(keyword_results[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - if not keyword_results: - return vector_results[:limit] - elif not vector_results: - return keyword_results[:limit] - else: - merged = self._merge_hybrid_results( - vector=vector_results, - keyword=keyword_results, - vector_weight=vector_weight, - text_weight=text_weight, - ) - - # Log merged results - logger.info("\n=== Merged Hybrid Results ===") - for i, r in enumerate(merged[:10], 1): - snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet - logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}") - - return merged[:limit] - elif self.vector_enabled: - vector_results = await self.vector_search(query, limit, sources) - return vector_results - elif self.fts_enabled: - keyword_results = await self.keyword_search(query, limit, sources) - return keyword_results - else: - return [] - - @staticmethod - def _merge_hybrid_results( - vector: list[MemorySearchResult], - keyword: list[MemorySearchResult], - vector_weight: float, - text_weight: float, - ) -> list[MemorySearchResult]: - """Merge vector and keyword search results with weighted scoring.""" - merged: dict[str, MemorySearchResult] = {} - - # Process vector results - for result in vector: - result.score = result.score * vector_weight - merged[result.merge_key] = result - - # Process keyword results - for result in keyword: - key = result.merge_key - if key in merged: - merged[key].score += result.score * text_weight - else: - result.score = result.score * text_weight - merged[key] = result - - # Sort by score and return - results = list(merged.values()) - results.sort(key=lambda r: r.score, reverse=True) - return results - - async def clear_all(self): - """Clear all indexed data.""" - cursor = self.conn.cursor() - try: - cursor.execute("BEGIN") - - cursor.execute(f"DELETE FROM {self.files_table_name}") - cursor.execute(f"DELETE FROM {self.chunks_table_name}") - - if self.vector_enabled: - cursor.execute(f"DELETE FROM {self.vector_table_name}") - - if self.fts_enabled: - cursor.execute(f"DELETE FROM {self.fts_table_name}") - - cursor.execute("COMMIT") - except Exception as e: - cursor.execute("ROLLBACK") - logger.error(f"Failed to clear all data: {e}") - raise - finally: - cursor.close() - - async def close(self): - """Close database connection.""" - if self.conn: - self.conn.close() - self.conn = None - await super().close() diff --git a/reme/core/file_store/zvec_file_store.py b/reme/core/file_store/zvec_file_store.py deleted file mode 100644 index 3d162819..00000000 --- a/reme/core/file_store/zvec_file_store.py +++ /dev/null @@ -1,573 +0,0 @@ -"""Zvec storage backend for file store.""" - -from __future__ import annotations - -import json -import time -from pathlib import Path -from typing import Any - -from .base_file_store import BaseFileStore -from ..enumeration import MemorySource -from ..schema import FileMetadata, MemoryChunk, MemorySearchResult -from ..utils import get_logger - -logger = get_logger() - -_ZVEC_IMPORT_ERROR: Exception | None = None - -try: - import zvec # type: ignore[import-untyped] - from zvec import ( - CollectionOption, - CollectionSchema, - DataType, - Doc, - FieldSchema, - HnswIndexParam, - InvertIndexParam, - VectorQuery, - VectorSchema, - ) - from zvec.typing import MetricType -except Exception as e: - _ZVEC_IMPORT_ERROR = e - zvec = None # type: ignore[assignment] - - -# zvec max topk (will be lifted to 100,000 in zvec v0.3.2+) -_ZVEC_MAX_TOPK = 1024 - -# Default vector field name -_DEFAULT_VECTOR_FIELD = "embedding" - - -def _escape(value: str) -> str: - """Escape a string value for zvec filter expressions.""" - return value.replace("'", "\\'") - - -def _build_file_store_schema(name: str, dimension: int) -> CollectionSchema: - """Build a zvec CollectionSchema for file store chunks.""" - return CollectionSchema( - name=name, - fields=[ - FieldSchema("content", DataType.STRING, nullable=True, index_param=InvertIndexParam()), - FieldSchema("path", DataType.STRING, nullable=True, index_param=InvertIndexParam()), - FieldSchema("source", DataType.STRING, nullable=True, index_param=InvertIndexParam()), - FieldSchema("start_line", DataType.INT64, nullable=True), - FieldSchema("end_line", DataType.INT64, nullable=True), - FieldSchema("hash", DataType.STRING, nullable=True), - FieldSchema("updated_at", DataType.INT64, nullable=True), - FieldSchema("file_metadata", DataType.STRING, nullable=True), - ], - vectors=[ - VectorSchema( - name=_DEFAULT_VECTOR_FIELD, - data_type=DataType.VECTOR_FP32, - dimension=dimension, - index_param=HnswIndexParam(metric_type=MetricType.COSINE), - ), - ], - ) - - -def _chunk_to_doc(chunk: MemoryChunk, file_meta_json: str = "{}") -> Doc: - """Convert a MemoryChunk to a zvec Doc.""" - fields: dict[str, Any] = { - "content": chunk.text, - "path": chunk.path, - "source": chunk.source.value if chunk.source else "", - "start_line": chunk.start_line, - "end_line": chunk.end_line, - "hash": chunk.hash, - "updated_at": int(time.time() * 1000), - "file_metadata": file_meta_json, - } - vectors: dict[str, Any] = {} - if chunk.embedding is not None: - vectors[_DEFAULT_VECTOR_FIELD] = chunk.embedding - return Doc(id=chunk.id, fields=fields, vectors=vectors) - - -def _doc_to_chunk(doc: Doc) -> MemoryChunk: - """Convert a zvec Doc to a MemoryChunk.""" - raw_vector = doc.vector(_DEFAULT_VECTOR_FIELD) - vector = raw_vector if isinstance(raw_vector, list) and len(raw_vector) > 0 else None - return MemoryChunk( - id=str(doc.id), - path=str(doc.field("path") or ""), - source=MemorySource(str(doc.field("source") or "")), - start_line=int(doc.field("start_line") or 0), - end_line=int(doc.field("end_line") or 0), - text=str(doc.field("content") or ""), - hash=str(doc.field("hash") or ""), - embedding=vector, - ) - - -def _build_source_filter(sources: list[MemorySource] | None) -> str | None: - """Build a zvec filter expression for source filtering.""" - if not sources: - return None - if len(sources) == 1: - return f"source='{_escape(sources[0].value)}'" - vals = ", ".join(f"'{_escape(s.value)}'" for s in sources) - return f"source IN ({vals})" - - -class ZvecFileStore(BaseFileStore): - """Zvec file storage with vector and keyword search. - - Provides zvec-backed persistent storage with: - - Vector similarity search (native zvec HNSW) - - Keyword search (Python substring matching on fetched results) - - Hybrid search (weighted fusion of vector and keyword results) - - Note: - Keyword search operates on chunks fetched from zvec, which is subject - to the topk limit (1024 in zvec < v0.3.2, 100,000 in v0.3.2+). - For collections with more chunks than the topk limit, keyword search - may not scan all documents. - """ - - def __init__( - self, - store_name: str, - db_path: str | Path, - embedding_model: Any | None = None, - vector_enabled: bool = False, - fts_enabled: bool = True, - dimension: int = 1024, - **kwargs: Any, - ): - if _ZVEC_IMPORT_ERROR is not None: - raise ImportError( - "Zvec requires extra dependencies. Install with `pip install zvec`", - ) from _ZVEC_IMPORT_ERROR - - super().__init__( - store_name=store_name, - db_path=db_path, - embedding_model=embedding_model, - vector_enabled=vector_enabled, - fts_enabled=fts_enabled, - **kwargs, - ) - - self.dimension = dimension - self._collection = None - self._initialized = False - self._metadata_file: Path = self.db_path / f"{store_name}_file_metadata.json" - self._metadata_cache: dict[str, dict[str, FileMetadata]] = {} - - @property - def collection_name(self) -> str: - """Get the name of the zvec collection for this store.""" - return f"chunks_{self.store_name}" - - # ------------------------------------------------------------------ - # Lifecycle - # ------------------------------------------------------------------ - - async def start(self) -> None: - """Initialize zvec engine and open the collection.""" - if not self._initialized: - try: - zvec.init() - except RuntimeError: - pass - self._initialized = True - - self.db_path.mkdir(parents=True, exist_ok=True) - collection_path = str(self.db_path / self.collection_name) - option = CollectionOption(read_only=False, enable_mmap=True) - - try: - self._collection = zvec.open(collection_path, option) - logger.info(f"Opened existing zvec file store collection: {collection_path}") - except Exception: - schema = _build_file_store_schema(self.collection_name, self.dimension) - self._collection = zvec.create_and_open( - path=collection_path, - schema=schema, - option=option, - ) - logger.info(f"Created new zvec file store collection: {collection_path}") - - self._metadata_cache = await self._load_metadata() - - async def close(self) -> None: - """Close zvec collection and persist metadata.""" - if self._metadata_cache: - await self._save_metadata(self._metadata_cache) - - if self._collection is not None: - try: - self._collection.flush() - except Exception as e: - logger.warning(f"Failed to flush collection on close: {e}") - self._collection = None - - # ------------------------------------------------------------------ - # Metadata management - # ------------------------------------------------------------------ - - async def _load_metadata(self) -> dict[str, dict[str, FileMetadata]]: - """Load file metadata from JSON file.""" - if not self._metadata_file.exists(): - return {} - try: - data = json.loads(self._metadata_file.read_text(encoding="utf-8")) - result: dict[str, dict[str, FileMetadata]] = {} - for source, files in data.items(): - result[source] = {} - for path, meta in files.items(): - result[source][path] = FileMetadata(**meta) - return result - except Exception as e: - logger.warning(f"Failed to load metadata from {self._metadata_file}: {e}") - return {} - - async def _save_metadata(self, metadata: dict[str, dict[str, FileMetadata]]) -> None: - """Save file metadata to JSON file.""" - try: - out: dict[str, dict[str, dict]] = {} - for source, files in metadata.items(): - out[source] = {} - for path, meta in files.items(): - out[source][path] = { - "path": meta.path, - "hash": meta.hash, - "mtime_ms": meta.mtime_ms, - "size": meta.size, - "chunk_count": meta.chunk_count, - } - self._metadata_file.write_text( - json.dumps(out, indent=2, ensure_ascii=False), - encoding="utf-8", - ) - except Exception as e: - logger.error(f"Failed to save metadata to {self._metadata_file}: {e}") - - # ------------------------------------------------------------------ - # CRUD operations - # ------------------------------------------------------------------ - - async def upsert_file( - self, - file_meta: FileMetadata, - source: MemorySource, - chunks: list[MemoryChunk], - ) -> None: - """Insert or update a file and its chunks.""" - if not chunks: - return - - # Delete existing chunks for this file first - await self.delete_file(file_meta.path, source) - - # Generate embeddings - chunks = await self.get_chunk_embeddings(chunks) - - file_meta_json = json.dumps( - { - "path": file_meta.path, - "hash": file_meta.hash, - "mtime_ms": file_meta.mtime_ms, - "size": file_meta.size, - "chunk_count": len(chunks), - }, - ensure_ascii=False, - ) - - docs = [_chunk_to_doc(c, file_meta_json) for c in chunks] - self._collection.insert(docs) - - # Update metadata cache - if source.value not in self._metadata_cache: - self._metadata_cache[source.value] = {} - self._metadata_cache[source.value][file_meta.path] = FileMetadata( - hash=file_meta.hash, - mtime_ms=file_meta.mtime_ms, - size=file_meta.size, - path=file_meta.path, - chunk_count=len(chunks), - ) - - async def delete_file(self, path: str, source: MemorySource) -> None: - """Delete a file and all its chunks.""" - filter_expr = f"path='{_escape(path)}' AND source='{_escape(source.value)}'" - results = self._collection.query(topk=_ZVEC_MAX_TOPK, filter=filter_expr, include_vector=False) - - ids_to_delete = [doc.id for doc in results] - if ids_to_delete: - self._collection.delete(ids_to_delete) - - if source.value in self._metadata_cache: - self._metadata_cache[source.value].pop(path, None) - - async def delete_file_chunks(self, path: str, chunk_ids: list[str]) -> None: - """Delete specific chunks for a file.""" - if not chunk_ids: - return - self._collection.delete(chunk_ids) - - async def upsert_chunks( - self, - chunks: list[MemoryChunk], - source: MemorySource, - ) -> None: - """Insert or update specific chunks.""" - if not chunks: - return - - chunks = await self.get_chunk_embeddings(chunks) - docs = [_chunk_to_doc(c) for c in chunks] - self._collection.upsert(docs) - - # ------------------------------------------------------------------ - # Listing and metadata - # ------------------------------------------------------------------ - - async def list_files(self, source: MemorySource) -> list[str]: - """List all indexed files for a source.""" - if source.value not in self._metadata_cache: - return [] - return list(self._metadata_cache[source.value].keys()) - - async def get_file_metadata( - self, - path: str, - source: MemorySource, - ) -> FileMetadata | None: - """Get file metadata.""" - if source.value not in self._metadata_cache: - return None - return self._metadata_cache[source.value].get(path) - - async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None: - """Update file metadata without affecting chunks.""" - if source.value not in self._metadata_cache: - self._metadata_cache[source.value] = {} - self._metadata_cache[source.value][file_meta.path] = FileMetadata( - hash=file_meta.hash, - mtime_ms=file_meta.mtime_ms, - size=file_meta.size, - path=file_meta.path, - chunk_count=file_meta.chunk_count, - ) - - async def get_file_chunks( - self, - path: str, - source: MemorySource, - ) -> list[MemoryChunk]: - """Get all chunks for a file.""" - filter_expr = f"path='{_escape(path)}' AND source='{_escape(source.value)}'" - results = self._collection.query( - topk=_ZVEC_MAX_TOPK, - filter=filter_expr, - include_vector=True, - ) - chunks = [_doc_to_chunk(doc) for doc in results] - chunks.sort(key=lambda c: c.start_line) - return chunks - - # ------------------------------------------------------------------ - # Search - # ------------------------------------------------------------------ - - async def vector_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform vector similarity search.""" - if not self.vector_enabled or not query: - return [] - - query_embedding = await self.get_embedding(query) - if not query_embedding: - return [] - - filter_expr = _build_source_filter(sources) - vq = VectorQuery(field_name=_DEFAULT_VECTOR_FIELD, vector=query_embedding) - - try: - results = self._collection.query( - vectors=vq, - topk=min(limit, _ZVEC_MAX_TOPK), - filter=filter_expr, - include_vector=False, - ) - except Exception as e: - logger.error(f"Vector search failed: {e}") - return [] - - search_results = [] - for doc in results: - score = doc.score if doc.score is not None else 0.0 - # zvec cosine score might need normalization depending on version - search_results.append( - MemorySearchResult( - path=str(doc.field("path") or ""), - start_line=int(doc.field("start_line") or 0), - end_line=int(doc.field("end_line") or 0), - score=score, - snippet=str(doc.field("content") or ""), - source=MemorySource(str(doc.field("source") or "")), - raw_metric=score, - ), - ) - - search_results.sort(key=lambda r: r.score, reverse=True) - return search_results[:limit] - - async def keyword_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - ) -> list[MemorySearchResult]: - """Perform keyword search via Python substring matching. - - Fetches chunks from zvec (subject to topk limit) then matches - keywords in Python. For collections larger than the topk limit, - not all documents are scanned. - """ - if not self.fts_enabled or not query: - return [] - - words = query.split() - if not words: - return [] - - # Fetch candidate chunks from zvec - filter_expr = _build_source_filter(sources) - results = self._collection.query( - topk=_ZVEC_MAX_TOPK, - filter=filter_expr, - include_vector=False, - ) - - query_lower = query.lower() - words_lower = [w.lower() for w in words] - n_words = len(words) - - search_results = [] - for doc in results: - text = str(doc.field("content") or "") - text_lower = text.lower() - match_count = sum(1 for w in words_lower if w in text_lower) - if match_count == 0: - continue - - base_score = match_count / n_words - phrase_bonus = 0.2 if n_words > 1 and query_lower in text_lower else 0.0 - score = min(1.0, base_score + phrase_bonus) - - search_results.append( - MemorySearchResult( - path=str(doc.field("path") or ""), - start_line=int(doc.field("start_line") or 0), - end_line=int(doc.field("end_line") or 0), - score=score, - snippet=text, - source=MemorySource(str(doc.field("source") or "")), - ), - ) - - search_results.sort(key=lambda r: r.score, reverse=True) - return search_results[:limit] - - async def hybrid_search( - self, - query: str, - limit: int, - sources: list[MemorySource] | None = None, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - ) -> list[MemorySearchResult]: - """Perform hybrid search combining vector and keyword search.""" - assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}" - - candidates = min(200, max(1, int(limit * candidate_multiplier))) - text_weight = 1.0 - vector_weight - - if self.vector_enabled and self.fts_enabled: - keyword_results = await self.keyword_search(query, candidates, sources) - vector_results = await self.vector_search(query, candidates, sources) - - if not keyword_results: - return vector_results[:limit] - elif not vector_results: - return keyword_results[:limit] - else: - return self._merge_hybrid_results( - vector=vector_results, - keyword=keyword_results, - vector_weight=vector_weight, - text_weight=text_weight, - )[:limit] - elif self.vector_enabled: - return await self.vector_search(query, limit, sources) - elif self.fts_enabled: - return await self.keyword_search(query, limit, sources) - else: - return [] - - @staticmethod - def _merge_hybrid_results( - vector: list[MemorySearchResult], - keyword: list[MemorySearchResult], - vector_weight: float, - text_weight: float, - ) -> list[MemorySearchResult]: - """Merge vector and keyword search results with weighted scoring.""" - merged: dict[str, MemorySearchResult] = {} - - for result in vector: - result.score = result.score * vector_weight - merged[result.merge_key] = result - - for result in keyword: - key = result.merge_key - if key in merged: - merged[key].score += result.score * text_weight - else: - result.score = result.score * text_weight - merged[key] = result - - results = list(merged.values()) - results.sort(key=lambda r: r.score, reverse=True) - return results - - # ------------------------------------------------------------------ - # Maintenance - # ------------------------------------------------------------------ - - async def clear_all(self) -> None: - """Clear all indexed data.""" - # Delete all documents - stats = self._collection.stats - count = stats.doc_count if stats else 0 - if count > 0: - try: - self._collection.delete_by_filter("content!=''") - except Exception: - remaining = count - while remaining > 0: - batch = self._collection.query( - topk=min(remaining, _ZVEC_MAX_TOPK), - include_vector=False, - ) - if not batch: - break - self._collection.delete([doc.id for doc in batch]) - remaining -= len(batch) - - self._metadata_cache = {} - await self._save_metadata({}) - logger.info(f"Cleared all data from zvec file store: {self.collection_name}") diff --git a/reme/core/file_watcher/__init__.py b/reme/core/file_watcher/__init__.py deleted file mode 100644 index 020d0cfa..00000000 --- a/reme/core/file_watcher/__init__.py +++ /dev/null @@ -1,19 +0,0 @@ -"""File watcher module for monitoring file system changes. - -This module provides file watcher implementations for monitoring file changes -and updating memory stores accordingly. -""" - -from .base_file_watcher import BaseFileWatcher -from .delta_file_watcher import DeltaFileWatcher -from .full_file_watcher import FullFileWatcher -from ..registry_factory import R - -__all__ = [ - "BaseFileWatcher", - "DeltaFileWatcher", - "FullFileWatcher", -] - -R.file_watchers.register("full")(FullFileWatcher) -R.file_watchers.register("delta")(DeltaFileWatcher) diff --git a/reme/core/file_watcher/base_file_watcher.py b/reme/core/file_watcher/base_file_watcher.py deleted file mode 100644 index 09b2f703..00000000 --- a/reme/core/file_watcher/base_file_watcher.py +++ /dev/null @@ -1,243 +0,0 @@ -"""Base file watcher implementation. - -This module provides the base class for file watcher implementations -that monitor file system changes and trigger callbacks. -""" - -import asyncio -from collections.abc import Coroutine -from pathlib import Path -from typing import Any, Callable - -from loguru import logger -from watchfiles import awatch, Change - -from ..enumeration import MemorySource -from ..file_store import BaseFileStore - - -class BaseFileWatcher: - """ - Minimal file watcher base class - - This base class provides basic file monitoring functionality that can be extended - to implement specific file monitoring requirements. - """ - - def __init__( - self, - watch_paths: list[str] | str, - suffix_filters: list[str] | None = None, - recursive: bool = False, - debounce: int = 2000, - chunk_tokens: int = 400, - chunk_overlap: int = 80, - file_store: BaseFileStore | None = None, - callback: Callable[[set[tuple[Change, str]]], None | Coroutine[Any, Any, None]] | None = None, - rebuild_index_on_start: bool = True, - poll_delay_ms: int = 2000, - **kwargs, - ): - """ - Initialize the file watcher - - Args: - watch_paths: Paths to watch for changes - suffix_filters: File suffix filters (e.g., ['.py', '.txt']) - recursive: Whether to watch directories recursively - debounce: Debounce time in milliseconds - chunk_tokens: Token size for chunking - chunk_overlap: Overlap size for chunks - file_store: File store instance - callback: Callback function for changes - rebuild_index_on_start: If True, clear all indexed data on start and rescan existing files. - If False, only monitor new changes without initialization. - poll_delay_ms: Polling delay in milliseconds. If > 300ms, force_polling will be enabled automatically. - **kwargs: Additional keyword arguments - """ - self.watch_paths: list[str] = [watch_paths] if isinstance(watch_paths, str) else watch_paths - self.suffix_filters: list[str] = suffix_filters or [] - self.recursive: bool = recursive - self.debounce: int = debounce - self.chunk_tokens: int = chunk_tokens - self.chunk_overlap: int = chunk_overlap - self.file_store: BaseFileStore = file_store - self.callback = callback - self.rebuild_index_on_start: bool = rebuild_index_on_start - self.poll_delay_ms: int = poll_delay_ms - self.kwargs: dict = kwargs - - self._stop_event = asyncio.Event() - self._watch_task: asyncio.Task | None = None - self._running = False - - async def start(self): - """Start the file watcher""" - if self._running: - return - - self._stop_event = asyncio.Event() - self._running = True - - async def _initialize_and_watch(): - if self.rebuild_index_on_start: - if self.file_store is not None: - await self.file_store.clear_all() - logger.info("Cleared all indexed data on start") - await self._scan_existing_files() - await self._watch_loop() - - self._watch_task = asyncio.create_task(_initialize_and_watch()) - logger.info(f"Started watching: {self.watch_paths}") - - async def close(self): - """Stop the file watcher""" - if not self._running: - return - - self._stop_event.set() - if self._watch_task: - await self._watch_task - self._running = False - logger.info("Stopped watching") - - def watch_filter(self, _change: Change, path: str) -> bool: - """Filter function for file watching.""" - # If no suffix filters are specified, watch all files - if not self.suffix_filters: - return True - - # Check if the file has one of the allowed suffixes - for suffix in self.suffix_filters: - if path.endswith("." + suffix.strip(".")): - return True - - return False - - async def _scan_existing_files(self): - """Scan existing files matching watch criteria and trigger on_changes with Change.added""" - existing_files: set[tuple[Change, str]] = set() - - for watch_path_str in self.watch_paths: - watch_path = Path(watch_path_str) - - if not watch_path.exists(): - logger.warning(f"Watch path does not exist: {watch_path}") - continue - - if watch_path.is_file(): - # Single file - if self.watch_filter(Change.added, str(watch_path)): - existing_files.add((Change.added, str(watch_path))) - elif watch_path.is_dir(): - # Directory - if self.recursive: - # Recursive scan - for file_path in watch_path.rglob("*"): - if file_path.is_file() and self.watch_filter(Change.added, str(file_path)): - existing_files.add((Change.added, str(file_path))) - else: - # Non-recursive scan (only immediate children) - for file_path in watch_path.iterdir(): - if file_path.is_file() and self.watch_filter(Change.added, str(file_path)): - existing_files.add((Change.added, str(file_path))) - - if existing_files: - logger.info(f"[SCAN_ON_START] Found {len(existing_files)} existing files matching watch criteria") - await self.on_changes(existing_files) - logger.info(f"[SCAN_ON_START] Added {len(existing_files)} files to memory store") - else: - logger.info("[SCAN_ON_START] No existing files found matching watch criteria") - - if self.file_store is not None: - files: list[str] = await self.file_store.list_files(MemorySource.MEMORY) - for file_path in files: - chunks = await self.file_store.get_file_chunks(file_path, MemorySource.MEMORY) - logger.info(f"Found existing file: {file_path}, {len(chunks)} chunks") - - async def _interruptible_sleep(self, seconds: float): - """Sleep that can be interrupted by stop_event.""" - try: - await asyncio.wait_for(self._stop_event.wait(), timeout=seconds) - except asyncio.TimeoutError: - pass # Normal timeout, continue - - async def _watch_loop(self): - """Core monitoring loop with auto-restart on failure""" - if not self.watch_paths: - logger.warning("No watch paths specified") - return - - while not self._stop_event.is_set(): - # Filter out non-existent paths before each watch attempt - valid_paths = [p for p in self.watch_paths if Path(p).exists()] - - if not valid_paths: - logger.warning("No valid watch paths exist, waiting 10 seconds before retry...") - await self._interruptible_sleep(10) - continue - - invalid_paths = set(self.watch_paths) - set(valid_paths) - if invalid_paths: - logger.warning(f"Skipping non-existent paths: {invalid_paths}") - - try: - logger.info(f"Starting watch on valid paths: {valid_paths}") - async for changes in awatch( - *valid_paths, - force_polling=True, - watch_filter=self.watch_filter, - recursive=self.recursive, - debounce=self.debounce, - poll_delay_ms=self.poll_delay_ms, - stop_event=self._stop_event, - ): - if self._stop_event.is_set(): - break - - await self.on_changes(changes) - - except FileNotFoundError as e: - # Watch path was deleted during monitoring - logger.error(f"Watch path no longer exists: {e}, restarting in 10 seconds...") - if not self._stop_event.is_set(): - await self._interruptible_sleep(10) - - except Exception as e: - # Log other exceptions and restart - logger.error(f"Error in watch loop: {e}, restarting in 10 seconds...", exc_info=True) - if not self._stop_event.is_set(): - await self._interruptible_sleep(10) - - async def _on_changes(self, changes: set[tuple[Change, str]]): - """Callback method to handle file changes""" - - async def on_changes(self, changes: set[tuple[Change, str]]): - """Hook method to handle file changes""" - if self.callback: - result = self.callback(changes) - if asyncio.iscoroutine(result): - await result - else: - await self._on_changes(changes) - logger.info(f"[{self.__class__.__name__}] on_changes: {changes}") - - def is_running(self) -> bool: - """Check if the watcher is running""" - return self._running - - async def add_path(self, path: str): - """Dynamically add a path to monitor""" - if path not in self.watch_paths: - self.watch_paths.append(path) - if self._running: - await self.close() - await self.start() - - async def remove_path(self, path: str): - """Remove a monitored path""" - if path in self.watch_paths: - self.watch_paths.remove(path) - if self._running: - await self.close() - await self.start() diff --git a/reme/core/file_watcher/delta_file_watcher.py b/reme/core/file_watcher/delta_file_watcher.py deleted file mode 100644 index 6148bd07..00000000 --- a/reme/core/file_watcher/delta_file_watcher.py +++ /dev/null @@ -1,280 +0,0 @@ -"""Delta file watcher for incremental file synchronization. - -This module provides a file watcher that detects append-only changes -and only processes newly added content, avoiding redundant operations. -""" - -import asyncio -import os - -from loguru import logger -from watchfiles import Change - -from .base_file_watcher import BaseFileWatcher -from ..enumeration import MemorySource -from ..schema import FileMetadata, MemoryChunk -from ..utils import chunk_markdown, hash_text - - -class DeltaFileWatcher(BaseFileWatcher): - """Delta file watcher implementation for incremental synchronization. - - This watcher detects append-only changes (e.g., log files) and only processes - the newly added content, avoiding redundant embedding requests for unchanged content. - - Strategy: - - Detect if file is append-only (new lines added at end) - - Find the safe cutoff point (considering chunk overlap) - - Only re-chunk and embed content from cutoff to end - - Delete affected old chunks and insert new chunks - """ - - def __init__(self, overlap_lines: int = 2, **kwargs): - """ - Initialize delta file watcher. - - Args: - chunk_tokens: Maximum tokens per chunk - chunk_overlap: Overlap tokens between chunks - """ - super().__init__(**kwargs) - self.overlap_lines = overlap_lines - self.dirty = False - - @staticmethod - async def _build_file_metadata(path: str) -> FileMetadata: - """Build file metadata from filesystem.""" - - def _read_file_sync(): - stat_t = os.stat(path) - with open(path, "r", encoding="utf-8") as f: - content_t = f.read() - return stat_t, content_t - - stat, content = await asyncio.to_thread(_read_file_sync) - return FileMetadata( - hash=hash_text(content), - mtime_ms=stat.st_mtime * 1000, - size=stat.st_size, - path=path, - content=content, - ) - - def _find_cutoff_line( - self, - old_chunks: list[MemoryChunk], - old_file_meta: FileMetadata, - new_file_meta: FileMetadata, - ) -> int | None: - """Find the safe cutoff line for incremental update. - - Uses a heuristic approach: if file size increased and hash changed, - we verify by comparing content. For true append-only files (like logs), - the old content should be a prefix of new content. - - Args: - old_chunks: Existing chunks sorted by start_line - old_file_meta: Previous file metadata - new_file_meta: Current file metadata (with content) - - Returns: - Cutoff line number (1-indexed), or None if not append-only - """ - if not old_chunks: - return None - - # File shrunk - definitely not append-only - if new_file_meta.size < old_file_meta.size: - logger.debug("File shrunk, not append-only") - return None - - # File didn't grow much - might be a modification - size_growth = new_file_meta.size - old_file_meta.size - if size_growth < 10: # Less than 10 bytes growth - logger.debug("Minimal size growth, treating as modification") - return None - - # Verify append-only by checking if old content is prefix - # We need to read old file content from chunks - old_chunks_sorted = sorted(old_chunks, key=lambda c: c.start_line) - - # Simple heuristic: check if first few chunks' content matches - # This avoids reconstructing full old content - new_lines = new_file_meta.content.split("\n") - - # Sample check: verify first chunk still matches - first_chunk = old_chunks_sorted[0] - first_chunk_lines = first_chunk.text.split("\n") - new_first_lines = new_lines[first_chunk.start_line - 1 : first_chunk.end_line] - - # Compare (allowing for minor whitespace differences at boundaries) - if len(first_chunk_lines) > 0 and len(new_first_lines) > 0: - # Check if most of the lines match - matches = sum(1 for old, new in zip(first_chunk_lines, new_first_lines) if old == new) - if matches < len(first_chunk_lines) * 0.8: # Less than 80% match - logger.debug("First chunk content changed, not append-only") - return None - - # File appears to be append-only - # Find the last chunk and set cutoff considering overlap - last_chunk = max(old_chunks_sorted, key=lambda c: c.end_line) - cutoff_line = max(1, last_chunk.end_line - self.overlap_lines) - - logger.debug( - f"Append-only detected: size {old_file_meta.size} -> {new_file_meta.size}, " - f"cutoff at line {cutoff_line}", - ) - - return cutoff_line - - @staticmethod - def _extract_content_from_line(content: str, start_line: int) -> str: - """Extract content starting from a specific line number.""" - lines = content.split("\n") - if start_line <= 1: - return content - if start_line > len(lines): - return "" - # start_line is 1-indexed, array is 0-indexed - return "\n".join(lines[start_line - 1 :]) - - async def _on_changes(self, changes: set[tuple[Change, str]]): - """Handle file changes with incremental synchronization.""" - self.dirty = True - - for change_type, path in changes: - if change_type == Change.added: - # New file: process everything - file_meta = await self._build_file_metadata(path) - chunks = ( - chunk_markdown( - file_meta.content, - file_meta.path, - MemorySource.MEMORY, - self.chunk_tokens, - self.chunk_overlap, - ) - or [] - ) - - if chunks: - chunks = await self.file_store.get_chunk_embeddings(chunks) - file_meta.chunk_count = len(chunks) - await self.file_store.upsert_file(file_meta, MemorySource.MEMORY, chunks) - logger.info(f"File added: {path} ({len(chunks)} chunks)") - else: - logger.warning(f"No chunks generated for new file {path}") - - elif change_type == Change.modified: - # Get existing data - old_chunks = await self.file_store.get_file_chunks(path, MemorySource.MEMORY) - old_file_meta = await self.file_store.get_file_metadata(path, MemorySource.MEMORY) - - # Read new file - file_meta = await self._build_file_metadata(path) - - # If no old chunks, fallback to full update - if not old_chunks or not old_file_meta: - logger.debug(f"No existing chunks for {path}, doing full update") - chunks = ( - chunk_markdown( - file_meta.content, - file_meta.path, - MemorySource.MEMORY, - self.chunk_tokens, - self.chunk_overlap, - ) - or [] - ) - if chunks: - chunks = await self.file_store.get_chunk_embeddings(chunks) - file_meta.chunk_count = len(chunks) - await self.file_store.delete_file(path, MemorySource.MEMORY) - await self.file_store.upsert_file( - file_meta, - MemorySource.MEMORY, - chunks, - ) - logger.info(f"File modified (full): {path} ({len(chunks)} chunks)") - continue - - # Check if append-only and find cutoff line - old_chunks_sorted = sorted(old_chunks, key=lambda c: c.start_line) - cutoff_line = self._find_cutoff_line(old_chunks_sorted, old_file_meta, file_meta) - - if cutoff_line is None: - # Not append-only, do full update - logger.debug(f"File {path} has modifications, doing full update") - chunks = ( - chunk_markdown( - file_meta.content, - file_meta.path, - MemorySource.MEMORY, - self.chunk_tokens, - self.chunk_overlap, - ) - or [] - ) - if chunks: - chunks = await self.file_store.get_chunk_embeddings(chunks) - file_meta.chunk_count = len(chunks) - await self.file_store.delete_file(path, MemorySource.MEMORY) - await self.file_store.upsert_file(file_meta, MemorySource.MEMORY, chunks) - logger.info(f"File modified (full): {path} ({len(chunks)} chunks)") - else: - # Append-only: incremental update - new_content_part = self._extract_content_from_line(file_meta.content, cutoff_line) - - new_chunks = ( - chunk_markdown( - new_content_part, - file_meta.path, - MemorySource.MEMORY, - self.chunk_tokens, - self.chunk_overlap, - ) - or [] - ) - - if not new_chunks: - logger.debug(f"No new chunks for {path}, skipping") - continue - - for idx, chunk in enumerate(new_chunks): - chunk.start_line += cutoff_line - 1 - chunk.end_line += cutoff_line - 1 - chunk.id = hash_text( - f"{chunk.source}:{chunk.path}:{chunk.start_line}:" f"{chunk.end_line}:{chunk.hash}:{idx}", - ) - - new_chunks = await self.file_store.get_chunk_embeddings(new_chunks) - - chunks_to_delete = [c.id for c in old_chunks_sorted if c.start_line >= cutoff_line] - - # Apply incremental updates - if chunks_to_delete: - await self.file_store.delete_file_chunks(path, chunks_to_delete) - - if new_chunks: - await self.file_store.upsert_chunks(new_chunks, MemorySource.MEMORY) - - # Update file metadata to reflect the changes - # Calculate new chunk count: old chunks - deleted + new chunks - new_chunk_count = len(old_chunks) - len(chunks_to_delete) + len(new_chunks) - file_meta.chunk_count = new_chunk_count - await self.file_store.update_file_metadata(file_meta, MemorySource.MEMORY) - - logger.info( - f"File modified (incremental): {path} " - f"(cutoff: line {cutoff_line}, " - f"+{len(new_chunks)} chunks, -{len(chunks_to_delete)} chunks)", - ) - - elif change_type == Change.deleted: - await self.file_store.delete_file(path, MemorySource.MEMORY) - logger.info(f"File deleted: {path}") - - else: - logger.warning(f"Unknown change type: {change_type}") - - self.dirty = False diff --git a/reme/core/file_watcher/full_file_watcher.py b/reme/core/file_watcher/full_file_watcher.py deleted file mode 100644 index 5b1852b9..00000000 --- a/reme/core/file_watcher/full_file_watcher.py +++ /dev/null @@ -1,79 +0,0 @@ -"""Full file watcher for complete file synchronization. - -This module provides a file watcher that processes entire files -on any change, ensuring complete synchronization. -""" - -import asyncio -from pathlib import Path - -from loguru import logger -from watchfiles import Change - -from .base_file_watcher import BaseFileWatcher -from ..enumeration import MemorySource -from ..schema import FileMetadata -from ..utils import chunk_markdown, hash_text - - -class FullFileWatcher(BaseFileWatcher): - """Full file watcher implementation for full synchronization""" - - def __init__(self, **kwargs): - """ - Initialize full file watcher""" - super().__init__(**kwargs) - self.dirty = False - - @staticmethod - async def _build_file_metadata(path: str) -> FileMetadata: - file_path = Path(path) - - def _read_file_sync(): - return file_path.stat(), file_path.read_text(encoding="utf-8") - - stat, content = await asyncio.to_thread(_read_file_sync) - return FileMetadata( - hash=hash_text(content), - mtime_ms=stat.st_mtime * 1000, - size=stat.st_size, - path=str(file_path.absolute()), - content=content, - ) - - async def _on_changes(self, changes: set[tuple[Change, str]]): - """Handle file changes with full synchronization""" - self.dirty = True - - for change_type, path in changes: - if change_type in [Change.added, Change.modified]: - file_meta = await self._build_file_metadata(path) - chunks = ( - chunk_markdown( - file_meta.content, - file_meta.path, - MemorySource.MEMORY, - self.chunk_tokens, - self.chunk_overlap, - ) - or [] - ) - if chunks: - chunks = await self.file_store.get_chunk_embeddings(chunks) - file_meta.chunk_count = len(chunks) - - await self.file_store.delete_file(file_meta.path, MemorySource.MEMORY) - logger.info(f"delete_file {file_meta.path}") - - await self.file_store.upsert_file(file_meta, MemorySource.MEMORY, chunks) - logger.info(f"Upserted {file_meta.chunk_count} chunks for {file_meta.path}") - - elif change_type == Change.deleted: - await self.file_store.delete_file(path, MemorySource.MEMORY) - logger.info(f"Deleted {path}") - - else: - logger.warning(f"Unknown change type: {change_type}") - - logger.info(f"File {change_type} changed: {path}") - self.dirty = False diff --git a/reme/core/flow/__init__.py b/reme/core/flow/__init__.py deleted file mode 100644 index 3b7e2753..00000000 --- a/reme/core/flow/__init__.py +++ /dev/null @@ -1,14 +0,0 @@ -"""flow""" - -from .base_flow import BaseFlow -from .cmd_flow import CmdFlow -from .expression_flow import ExpressionFlow -from ..registry_factory import R - -__all__ = [ - "BaseFlow", - "CmdFlow", - "ExpressionFlow", -] - -R.flows.register(ExpressionFlow) diff --git a/reme/core/flow/base_flow.py b/reme/core/flow/base_flow.py deleted file mode 100644 index 68e66fd5..00000000 --- a/reme/core/flow/base_flow.py +++ /dev/null @@ -1,208 +0,0 @@ -"""Base flow module providing abstract flow execution with caching and operation orchestration.""" - -import asyncio -import hashlib -import json -from abc import ABC, abstractmethod - -from loguru import logger - -from ..enumeration import ChunkEnum -from ..op import BaseOp, SequentialOp, ParallelOp -from ..registry_factory import R -from ..runtime_context import RuntimeContext -from ..schema import Response, ToolCall -from ..service_context import ServiceContext -from ..utils import camel_to_snake, CacheHandler - - -class BaseFlow(ABC): - """Abstract base class for flow execution with caching, streaming, and operation tree management.""" - - def __init__( - self, - name: str = "", - stream: bool = False, - raise_exception: bool = True, - enable_cache: bool = False, - cache_path: str = "cache/flow", - cache_expire_hours: float = 0.1, - service_context: ServiceContext | None = None, - **kwargs, - ): - """Initialize flow configuration and execution state.""" - super().__init__() - - self.name: str = name or camel_to_snake(self.__class__.__name__) - self.stream: bool = stream - self.raise_exception: bool = raise_exception - self.enable_cache: bool = enable_cache - self.cache_path: str = cache_path - self.cache_expire_hours: float = cache_expire_hours - self.service_context: ServiceContext | None = service_context - self.flow_params: dict = kwargs - - self._cache: CacheHandler | None = None - self._flow_printed: bool = False - self._flow_op: BaseOp | None = None - self._tool_call: ToolCall | None = None - - def _build_tool_call(self) -> ToolCall | None: - """Generate the tool call schema definition for this flow.""" - - @abstractmethod - def _build_flow(self) -> BaseOp: - """Construct the root operation tree for flow execution.""" - - def _compute_cache_key(self, params: dict) -> str | None: - """Generate a SHA256 hash from input parameters for caching.""" - try: - payload = json.dumps(params, sort_keys=True, ensure_ascii=False, default=str) - return hashlib.sha256(payload.encode("utf-8")).hexdigest() - except Exception as e: - logger.exception(f"[{self.__class__.__name__}] {self.name} cache key serialization failed: {e}") - return None - - def _maybe_load_cached(self, params: dict) -> Response | None: - """Retrieve a cached response if caching is enabled and available.""" - if not self.enable_cache or self.stream: - return None - - if key := self._compute_cache_key(params): - if cached := self.cache.load(key): - logger.info(f"[{self.__class__.__name__}] Loaded {self.name} response from cache.") - return Response(**cached) - return None - - def _maybe_save_cache(self, params: dict, response: Response): - """Persist the execution response to the cache.""" - if not self.enable_cache or self.stream: - return - - if key := self._compute_cache_key(params): - self.cache.save(key, response.model_dump(exclude_none=True), expire_hours=self.cache_expire_hours) - - def _print_operation_tree(self, name: str, op: BaseOp, indent: int): - """Recursively log the hierarchy of the flow's operation tree.""" - prefix = " " * indent - op_type = "sequential" if isinstance(op, SequentialOp) else "parallel" if isinstance(op, ParallelOp) else name - logger.info(f"[{self.__class__.__name__}] {prefix}{op_type} execution") - - for sub_op in op.sub_ops or []: - self._print_operation_tree(sub_op.name, sub_op, indent + 2) - - @property - def tool_call(self) -> ToolCall | None: - """Lazily construct the ToolCall schema describing this flow.""" - if hasattr(self.flow_op, "tool_call"): - return self.flow_op.tool_call - - if self._tool_call is None: - self._tool_call = self._build_tool_call() - if self._tool_call: - self._tool_call.name = self._tool_call.name or self.name - return self._tool_call - - @property - def cache(self) -> CacheHandler: - """Provide access to the internal CacheHandler instance.""" - assert self.enable_cache, "Cache usage requested while disabled." - if self._cache is None: - self._cache = CacheHandler(f"{self.cache_path}/{self.name}") - return self._cache - - @property - def flow_op(self) -> BaseOp: - """Lazily build and retrieve the root operation of the flow.""" - if self._flow_op is None: - self._flow_op = self._build_flow() - return self._flow_op - - @property - def async_mode(self) -> bool: - """Check if the current flow operation tree is asynchronous.""" - return self.flow_op.async_mode - - @staticmethod - def parse_expression(expression: str) -> BaseOp: - """Parse a string expression into an executable BaseOp instance.""" - lines = [x.strip() for x in expression.strip().splitlines() if x.strip()] - if not lines: - raise ValueError("Expression is empty") - - if len(lines) > 1: - exec("\n".join(lines[:-1]), {"__builtins__": {}}, R.ops) - - result = eval(lines[-1], {"__builtins__": {}}, R.ops) - if not isinstance(result, BaseOp): - raise TypeError(f"Expression evaluated to {type(result)}, expected BaseOp") - return result - - def print_flow(self): - """Log the visual structure of the flow once.""" - if not self._flow_printed: - logger.info(f"[{self.__class__.__name__}] ---------- [Flow Structure] {self.name} [Start] ----------") - self._print_operation_tree(self.name, self.flow_op, 0) - logger.info(f"[{self.__class__.__name__}] ---------- [Flow Structure] {self.name} [End] ----------") - self._flow_printed = True - - async def call(self, **kwargs) -> Response | asyncio.Queue: - """Execute the flow asynchronously with parameter caching.""" - kwargs["stream"] = self.stream - logger.info(f"[{self.__class__.__name__}] {self.name} incoming params: {kwargs}") - if cached := self._maybe_load_cached(kwargs): - return cached - - context = RuntimeContext(service_context=self.service_context, **kwargs) - try: - self.print_flow() - flow_op: BaseOp = self._build_flow() - assert self.flow_op.async_mode, "Async call requires an async flow operation." - await flow_op.call(context=context) - - if self.stream: - await context.add_stream_done() - return context.stream_queue - - else: - self._maybe_save_cache(kwargs, context.response) - return context.response - - except Exception as e: - logger.exception(f"[{self.__class__.__name__}] {self.name} async call failed: {e}") - if self.raise_exception: - raise e - - if self.stream: - await context.add_stream_chunk_and_type(str(e), ChunkEnum.ERROR) - await context.add_stream_done() - return context.stream_queue - - else: - context.add_response_error(e) - return context.response - - def call_sync(self, **kwargs) -> Response: - """Execute the flow synchronously with parameter caching.""" - logger.info(f"[{self.__class__.__name__}] {self.name} incoming sync params: {kwargs}") - assert not self.stream, "Synchronous call cannot be used in stream mode." - if cached := self._maybe_load_cached(kwargs): - return cached - - context = RuntimeContext(service_context=self.service_context, **kwargs) - try: - self.print_flow() - flow_op: BaseOp = self._build_flow() - assert not self.flow_op.async_mode, "Sync call requires a sync flow operation." - flow_op.call_sync(context=context) - - self._maybe_save_cache(kwargs, context.response) - return context.response - - except Exception as e: - logger.exception(f"[{self.__class__.__name__}] {self.name} sync call failed: {e}") - if self.raise_exception: - raise e - - context.add_response_error(e) - return context.response diff --git a/reme/core/flow/cmd_flow.py b/reme/core/flow/cmd_flow.py deleted file mode 100644 index 695ab48e..00000000 --- a/reme/core/flow/cmd_flow.py +++ /dev/null @@ -1,18 +0,0 @@ -"""Command-based flow implementation for parsing and executing operation sequences.""" - -from .base_flow import BaseFlow -from ..op import BaseOp - - -class CmdFlow(BaseFlow): - """A flow class that builds an operation chain from a string expression.""" - - def __init__(self, flow: str = "", **kwargs): - """Initialize the command flow with a string-based operation definition.""" - super().__init__(**kwargs) - self.flow = flow - assert flow, "add `cmd.flow=` in cmd!" - - def _build_flow(self) -> BaseOp: - """Parse the stored flow expression into a functional operation object.""" - return self.parse_expression(self.flow) diff --git a/reme/core/flow/expression_flow.py b/reme/core/flow/expression_flow.py deleted file mode 100644 index a844c719..00000000 --- a/reme/core/flow/expression_flow.py +++ /dev/null @@ -1,37 +0,0 @@ -"""Expression-based flow implementation driven by configuration objects.""" - -from .base_flow import BaseFlow -from ..op import BaseOp -from ..schema import FlowConfig, ToolCall -from ..service_context import ServiceContext - - -class ExpressionFlow(BaseFlow): - """A flow implementation that constructs operations from a FlowConfig definition.""" - - def __init__(self, flow_config: FlowConfig, service_context: ServiceContext): - """Initialize the flow using settings and metadata from a FlowConfig instance.""" - self.flow_config: FlowConfig = flow_config - super().__init__( - name=flow_config.name, - stream=self.flow_config.stream, - raise_exception=self.flow_config.raise_exception, - enable_cache=self.flow_config.enable_cache, - cache_path=self.flow_config.cache_path, - cache_expire_hours=self.flow_config.cache_expire_hours, - service_context=service_context, - **flow_config.model_extra, - ) - - def _build_flow(self) -> BaseOp: - """Generate the operation chain by parsing the flow content string.""" - return self.parse_expression(self.flow_config.flow_content) - - def _build_tool_call(self) -> ToolCall: - """Construct a tool call representation based on configuration parameters.""" - return ToolCall( - **{ - "description": self.flow_config.description, - "parameters": self.flow_config.parameters, - }, - ) diff --git a/reme/core/llm/__init__.py b/reme/core/llm/__init__.py deleted file mode 100644 index baa7808e..00000000 --- a/reme/core/llm/__init__.py +++ /dev/null @@ -1,21 +0,0 @@ -"""llm""" - -from .base_llm import BaseLLM -from .lite_llm import LiteLLM -from .lite_llm_sync import LiteLLMSync -from .openai_llm import OpenAILLM -from .openai_llm_sync import OpenAILLMSync -from ..registry_factory import R - -__all__ = [ - "BaseLLM", - "LiteLLM", - "LiteLLMSync", - "OpenAILLM", - "OpenAILLMSync", -] - -R.llms.register("litellm")(LiteLLM) -R.llms.register("litellm_sync")(LiteLLMSync) -R.llms.register("openai")(OpenAILLM) -R.llms.register("openai_sync")(OpenAILLMSync) diff --git a/reme/core/llm/base_llm.py b/reme/core/llm/base_llm.py deleted file mode 100644 index c08a423c..00000000 --- a/reme/core/llm/base_llm.py +++ /dev/null @@ -1,510 +0,0 @@ -"""Base interface for LLM implementations.""" - -import asyncio -import json -import time -from abc import ABC, abstractmethod -from typing import Callable, Generator, AsyncGenerator, Any - -from loguru import logger - -from ..enumeration import ChunkEnum, Role -from ..schema import Message, StreamChunk, ToolCall -from ..utils import extract_content - - -class BaseLLM(ABC): - """Base class for LLM interactions.""" - - def __init__( - self, - api_key: str | None = None, - base_url: str | None = None, - model_name: str = "", - max_retries: int = 10, - raise_exception: bool = False, - request_interval: float = 0.0, - **kwargs, - ): - """Initialize LLM client. - - Args: - model_name: Model name to use - max_retries: Maximum retry attempts on failure - raise_exception: Raise exceptions or return default values - request_interval: Minimum seconds between requests (default: 0.0) - **kwargs: Additional model-specific parameters - """ - self.api_key: str = api_key - self.base_url: str = base_url - self.model_name: str = model_name - self.max_retries: int = max_retries - self.raise_exception: bool = raise_exception - self.request_interval: float = request_interval - self.kwargs: dict = kwargs - - self._last_request_time: float = 0.0 - self._request_lock: asyncio.Lock = asyncio.Lock() - - @staticmethod - def _accumulate_tool_call_chunk(tool_call, ret_tools: list[ToolCall]): - """Assemble incremental tool call chunks into complete ToolCall objects.""" - index = tool_call.index - - while len(ret_tools) <= index: - ret_tools.append(ToolCall(index=index)) - - if tool_call.id: - ret_tools[index].id += tool_call.id - - if tool_call.function and tool_call.function.name: - ret_tools[index].name += tool_call.function.name - - if tool_call.function and tool_call.function.arguments: - ret_tools[index].arguments += tool_call.function.arguments - - @staticmethod - def _validate_and_serialize_tools(ret_tool_calls: list[ToolCall], tools: list[ToolCall]) -> list[dict]: - """Validate and serialize tool calls.""" - if not ret_tool_calls: - return [] - - tool_dict: dict[str, ToolCall] = {x.name: x for x in tools} if tools else {} - validated_tools = [] - - for tool in ret_tool_calls: - if tool.name not in tool_dict: - continue - - if not tool.sanitize_and_check_argument(): - logger.error(f"Invalid JSON arguments in {tool.name}: {tool.arguments}") - raise ValueError(f"Invalid JSON arguments in {tool.name}: {tool.arguments}") - - validated_tools.append(tool.simple_output_dump()) - return validated_tools - - @abstractmethod - def _build_stream_kwargs( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - log_params: bool = True, - model_name: str | None = None, - **kwargs, - ) -> dict: - """Build provider-specific streaming parameters.""" - - async def _stream_chat( - self, - messages: list[Message], - tools: list[ToolCall] | None, - stream_kwargs: dict, - ) -> AsyncGenerator[StreamChunk, None]: - """Async generator for streaming response chunks.""" - - def _stream_chat_sync( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - stream_kwargs: dict | None = None, - ) -> Generator[StreamChunk, None, None]: - """Sync generator for streaming response chunks.""" - - async def stream_chat( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - model_name: str | None = None, - **kwargs, - ) -> AsyncGenerator[StreamChunk, None]: - """Stream chat completions with retries and return final message.""" - if self.request_interval > 0: - async with self._request_lock: - current_time = time.time() - elapsed = current_time - self._last_request_time - if elapsed < self.request_interval: - await asyncio.sleep(self.request_interval - elapsed) - self._last_request_time = time.time() - - async for chunk in self._stream_chat_impl(messages, tools, model_name, **kwargs): - yield chunk - - async def _stream_chat_impl( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - model_name: str | None = None, - **kwargs, - ) -> AsyncGenerator[StreamChunk, None]: - """Stream chat with retry logic.""" - stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) - - for i in range(self.max_retries): - try: - async for chunk in self._stream_chat(messages=messages, tools=tools, stream_kwargs=stream_kwargs): - yield chunk - - break - - except Exception as e: - logger.exception(f"Stream chat error (model={self.model_name}): {e.args}") - - if i == self.max_retries - 1: - if self.raise_exception: - raise e - yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e)) - break - - yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e)) - await asyncio.sleep(i + 1) - - def stream_chat_sync( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - model_name: str | None = None, - **kwargs, - ) -> Generator[StreamChunk, None, None]: - """Stream chat completions synchronously with retries.""" - stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) - - for i in range(self.max_retries): - try: - yield from self._stream_chat_sync(messages=messages, tools=tools, stream_kwargs=stream_kwargs) - break - - except Exception as e: - logger.exception(f"Stream chat sync error (model={self.model_name}): {e.args}") - - if i == self.max_retries - 1: - if self.raise_exception: - raise e - yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e)) - break - - yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e)) - time.sleep(i + 1) - - async def _chat( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - enable_stream_print: bool = False, - model_name: str | None = None, - **kwargs, - ) -> Message: - """Aggregate full response by consuming the stream.""" - state = { - "enter_think": False, - "enter_answer": False, - "reasoning_content": "", - "answer_content": "", - "tool_calls": [], - } - - stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) - async for stream_chunk in self._stream_chat(messages=messages, tools=tools, stream_kwargs=stream_kwargs): - if stream_chunk.chunk_type is ChunkEnum.USAGE: - if enable_stream_print: - print( - f"\n{json.dumps(stream_chunk.chunk, ensure_ascii=False, indent=2)}", - flush=True, - ) - - elif stream_chunk.chunk_type is ChunkEnum.THINK: - if enable_stream_print: - if not state["enter_think"]: - state["enter_think"] = True - print("\n", end="", flush=True) - print(stream_chunk.chunk, end="", flush=True) - state["reasoning_content"] += stream_chunk.chunk - - elif stream_chunk.chunk_type is ChunkEnum.ANSWER: - if enable_stream_print: - if not state["enter_answer"]: - state["enter_answer"] = True - if state["enter_think"]: - print("\n", flush=True) - print(stream_chunk.chunk, end="", flush=True) - state["answer_content"] += stream_chunk.chunk - - elif stream_chunk.chunk_type is ChunkEnum.TOOL: - if enable_stream_print: - print(f"\n{json.dumps(stream_chunk.chunk, ensure_ascii=False, indent=2)}", flush=True) - state["tool_calls"].append(stream_chunk.chunk) - - elif stream_chunk.chunk_type is ChunkEnum.ERROR: - if enable_stream_print: - print(f"\n{stream_chunk.chunk}", flush=True) - - return Message( - role=Role.ASSISTANT, - reasoning_content=state["reasoning_content"], - content=state["answer_content"], - tool_calls=state["tool_calls"], - ) - - def _chat_sync( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - enable_stream_print: bool = False, - model_name: str | None = None, - **kwargs, - ) -> Message: - """Aggregate full response synchronously by consuming the stream.""" - state = { - "enter_think": False, - "enter_answer": False, - "reasoning_content": "", - "answer_content": "", - "tool_calls": [], - } - - stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) - for stream_chunk in self._stream_chat_sync(messages=messages, tools=tools, stream_kwargs=stream_kwargs): - if stream_chunk.chunk_type is ChunkEnum.USAGE: - if enable_stream_print: - print( - f"\n{json.dumps(stream_chunk.chunk, ensure_ascii=False, indent=2)}", - flush=True, - ) - - elif stream_chunk.chunk_type is ChunkEnum.THINK: - if enable_stream_print: - if not state["enter_think"]: - state["enter_think"] = True - print("\n", end="", flush=True) - print(stream_chunk.chunk, end="", flush=True) - state["reasoning_content"] += stream_chunk.chunk - - elif stream_chunk.chunk_type is ChunkEnum.ANSWER: - if enable_stream_print: - if not state["enter_answer"]: - state["enter_answer"] = True - if state["enter_think"]: - print("\n", flush=True) - print(stream_chunk.chunk, end="", flush=True) - state["answer_content"] += stream_chunk.chunk - - elif stream_chunk.chunk_type is ChunkEnum.TOOL: - if enable_stream_print: - print(f"\n{json.dumps(stream_chunk.chunk, ensure_ascii=False, indent=2)}", flush=True) - state["tool_calls"].append(stream_chunk.chunk) - - elif stream_chunk.chunk_type is ChunkEnum.ERROR: - if enable_stream_print: - print(f"\n{stream_chunk.chunk}", flush=True) - - return Message( - role=Role.ASSISTANT, - reasoning_content=state["reasoning_content"], - content=state["answer_content"], - tool_calls=state["tool_calls"], - ) - - async def chat( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - enable_stream_print: bool = False, - callback_fn: Callable[[Message], Any] | None = None, - default_value: Any = None, - model_name: str | None = None, - **kwargs, - ) -> Message | Any: - """Chat completion with retries and error handling.""" - if self.request_interval > 0: - async with self._request_lock: - current_time = time.time() - elapsed = current_time - self._last_request_time - if elapsed < self.request_interval: - await asyncio.sleep(self.request_interval - elapsed) - self._last_request_time = time.time() - - return await self._chat_impl( - messages, - tools, - enable_stream_print, - callback_fn, - default_value, - model_name, - **kwargs, - ) - - async def _chat_impl( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - enable_stream_print: bool = False, - callback_fn: Callable[[Message], Any] | None = None, - default_value: Any = None, - model_name: str | None = None, - **kwargs, - ) -> Message | Any: - """Chat with retry and error handling logic.""" - effective_model = model_name if model_name is not None else self.model_name - - for i in range(self.max_retries): - try: - result = await self._chat( - messages=messages, - tools=tools, - enable_stream_print=enable_stream_print, - model_name=model_name, - **kwargs, - ) - return callback_fn(result) if callback_fn else result - - except Exception as e: - error_message = str(e.args[0]) if e.args else str(e) - is_inappropriate_content = "inappropriate content" in error_message.lower() - is_rate_limit_error = ( - "request rate increased too quickly" in error_message.lower() - or "exceeded your current quota" in error_message.lower() - or "insufficient_quota" in error_message.lower() - ) - - if is_inappropriate_content: - logger.error(f"Inappropriate content detected (model={effective_model})") - logger.error("=" * 80) - for idx, msg in enumerate(messages): - logger.error(f"Message {idx + 1} [role={msg.role}]:") - logger.error(f"Content: {msg.content}") - if msg.reasoning_content: - logger.error(f"Reasoning: {msg.reasoning_content}") - if msg.tool_calls: - logger.error(f"Tool calls: {msg.tool_calls}") - logger.error("-" * 80) - logger.error("=" * 80) - return Message(role=Role.ASSISTANT, content="") - - if is_rate_limit_error: - logger.warning( - f"Rate limit hit (model={effective_model}), sleeping 60s (attempt {i + 1}/{self.max_retries})", - ) - await asyncio.sleep(60) - continue - - logger.exception(f"Chat error (model={effective_model}): {e.args}") - - if i == self.max_retries - 1: - if self.raise_exception: - raise e - return default_value - - await asyncio.sleep(1 + i) - return default_value - - def chat_sync( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - enable_stream_print: bool = False, - callback_fn: Callable[[Message], Any] | None = None, - default_value: Any = None, - model_name: str | None = None, - **kwargs, - ) -> Message | Any: - """Chat completion synchronously with retries and error handling.""" - effective_model = model_name if model_name is not None else self.model_name - - for i in range(self.max_retries): - try: - result = self._chat_sync( - messages=messages, - tools=tools, - enable_stream_print=enable_stream_print, - model_name=model_name, - **kwargs, - ) - return callback_fn(result) if callback_fn else result - - except Exception as e: - error_message = str(e.args[0]) if e.args else str(e) - is_inappropriate_content = "inappropriate content" in error_message.lower() - is_rate_limit_error = ( - "request rate increased too quickly" in error_message.lower() - or "exceeded your current quota" in error_message.lower() - or "insufficient_quota" in error_message.lower() - ) - - if is_inappropriate_content: - logger.error(f"Inappropriate content detected (model={effective_model})") - logger.error("=" * 80) - for idx, msg in enumerate(messages): - logger.error(f"Message {idx + 1} [role={msg.role}]:") - logger.error(f"Content: {msg.content}") - if msg.reasoning_content: - logger.error(f"Reasoning: {msg.reasoning_content}") - if msg.tool_calls: - logger.error(f"Tool calls: {msg.tool_calls}") - logger.error("-" * 80) - logger.error("=" * 80) - return Message(role=Role.ASSISTANT, content="") - - if is_rate_limit_error: - logger.warning( - f"Rate limit hit (model={effective_model}), sleeping 60s (attempt {i + 1}/{self.max_retries})", - ) - time.sleep(60) - continue - - logger.exception(f"Chat sync error (model={effective_model}): {e.args}") - - if i == self.max_retries - 1: - if self.raise_exception: - raise e - return default_value - - time.sleep(1 + i) - return default_value - - async def simple_request( - self, - prompt: str, - model_name: str | None = None, - callback_fn: Callable[[Message], Any] | None = None, - default_value: Any = None, - **kwargs, - ) -> str: - """Make a simple request using the LLM.""" - assistant_message = await self.chat( - messages=[Message(role=Role.USER, content=prompt)], - model_name=model_name, - callback_fn=callback_fn, - default_value=default_value, - **kwargs, - ) - return assistant_message.content - - async def simple_request_for_json( - self, - prompt: str, - model_name: str | None = None, - **kwargs, - ) -> dict: - """Make a simple request using the LLM and extract JSON.""" - - def extract_fn(message: Message) -> dict: - return extract_content(message.content) - - return await self.chat( - messages=[Message(role=Role.USER, content=prompt)], - model_name=model_name, - callback_fn=extract_fn, - default_value={}, - **kwargs, - ) - - def start_sync(self): - """Synchronously initialize resources.""" - - async def start(self): - """Asynchronously initialize resources.""" - - def close_sync(self): - """Synchronously release resources and close connections.""" - - async def close(self): - """Asynchronously release resources and close connections.""" diff --git a/reme/core/llm/lite_llm.py b/reme/core/llm/lite_llm.py deleted file mode 100644 index 13bc1872..00000000 --- a/reme/core/llm/lite_llm.py +++ /dev/null @@ -1,104 +0,0 @@ -"""LiteLLM asynchronous implementation for ReMe.""" - -from typing import AsyncGenerator - -from loguru import logger - -from .base_llm import BaseLLM -from ..enumeration import ChunkEnum -from ..schema import Message, StreamChunk, ToolCall - - -class LiteLLM(BaseLLM): - """Async LLM implementation using LiteLLM to support multiple providers.""" - - def __init__(self, custom_llm_provider: str = "openai", **kwargs): - """Initialize the LiteLLM client with API configuration and provider settings.""" - super().__init__(**kwargs) - self.custom_llm_provider: str = custom_llm_provider - - def _build_stream_kwargs( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - log_params: bool = True, - model_name: str | None = None, - **kwargs, - ) -> dict: - """Construct and log the parameters dictionary for LiteLLM API calls. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - log_params: Whether to log parameters - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ - # Use the provided model_name or fall back to self.model_name - effective_model = model_name if model_name is not None else self.model_name - - # Construct the API parameters by merging multiple sources - llm_kwargs = { - "model": effective_model, - "messages": [x.simple_dump() for x in messages], - "tools": [x.simple_input_dump() for x in tools] if tools else None, - "stream": True, - "custom_llm_provider": self.custom_llm_provider, - **self.kwargs, - **kwargs, - } - - # Add API key and base URL if provided - if self.api_key: - llm_kwargs["api_key"] = self.api_key - if self.base_url: - llm_kwargs["base_url"] = self.base_url - - # Log parameters for debugging, with message/tool counts instead of full content - if log_params: - log_kwargs: dict = {} - for k, v in llm_kwargs.items(): - if k in ["messages", "tools"]: - log_kwargs[k] = len(v) if v is not None else 0 - elif k == "api_key": - # Mask API key in logs for security - log_kwargs[k] = "***" if v else None - else: - log_kwargs[k] = v - logger.info(f"llm_kwargs={log_kwargs}") - - return llm_kwargs - - async def _stream_chat( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - stream_kwargs: dict | None = None, - ) -> AsyncGenerator[StreamChunk, None]: - """Execute async streaming chat requests and yield processed response chunks.""" - import litellm - - stream_kwargs = stream_kwargs or {} - completion = await litellm.acompletion(**stream_kwargs) - ret_tool_calls: list[ToolCall] = [] - - async for chunk in completion: - if not chunk.choices: - if hasattr(chunk, "usage") and chunk.usage: - yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump()) - continue - - delta = chunk.choices[0].delta - - if hasattr(delta, "reasoning_content") and delta.reasoning_content: - yield StreamChunk(chunk_type=ChunkEnum.THINK, chunk=delta.reasoning_content) - - if delta.content: - yield StreamChunk(chunk_type=ChunkEnum.ANSWER, chunk=delta.content) - - if hasattr(delta, "tool_calls") and delta.tool_calls is not None: - for tool_call in delta.tool_calls: - self._accumulate_tool_call_chunk(tool_call, ret_tool_calls) - - for tool_data in self._validate_and_serialize_tools(ret_tool_calls, tools): - yield StreamChunk(chunk_type=ChunkEnum.TOOL, chunk=tool_data) diff --git a/reme/core/llm/lite_llm_sync.py b/reme/core/llm/lite_llm_sync.py deleted file mode 100644 index d074a9f3..00000000 --- a/reme/core/llm/lite_llm_sync.py +++ /dev/null @@ -1,45 +0,0 @@ -"""Synchronous LiteLLM-based LLM implementation for the ReMe framework.""" - -from typing import Generator - -from .lite_llm import LiteLLM -from ..enumeration import ChunkEnum -from ..schema import Message, StreamChunk, ToolCall - - -class LiteLLMSync(LiteLLM): - """Synchronous LiteLLM client for executing chat completions and streaming responses.""" - - def _stream_chat_sync( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - stream_kwargs: dict | None = None, - ) -> Generator[StreamChunk, None, None]: - """Internal synchronous generator for processing streaming chat completion chunks.""" - import litellm - - stream_kwargs = stream_kwargs or {} - completion = litellm.completion(**stream_kwargs) - ret_tool_calls: list[ToolCall] = [] - - for chunk in completion: - if not chunk.choices: - if hasattr(chunk, "usage") and chunk.usage: - yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump()) - continue - - delta = chunk.choices[0].delta - - if hasattr(delta, "reasoning_content") and delta.reasoning_content: - yield StreamChunk(chunk_type=ChunkEnum.THINK, chunk=delta.reasoning_content) - - if delta.content: - yield StreamChunk(chunk_type=ChunkEnum.ANSWER, chunk=delta.content) - - if hasattr(delta, "tool_calls") and delta.tool_calls is not None: - for tool_call in delta.tool_calls: - self._accumulate_tool_call_chunk(tool_call, ret_tool_calls) - - for tool_data in self._validate_and_serialize_tools(ret_tool_calls, tools): - yield StreamChunk(chunk_type=ChunkEnum.TOOL, chunk=tool_data) diff --git a/reme/core/llm/openai_llm.py b/reme/core/llm/openai_llm.py deleted file mode 100644 index 6dd90887..00000000 --- a/reme/core/llm/openai_llm.py +++ /dev/null @@ -1,113 +0,0 @@ -"""Asynchronous OpenAI-compatible LLM implementation supporting streaming, tool calls, and reasoning content.""" - -from typing import AsyncGenerator - -from loguru import logger -from openai import AsyncOpenAI - -from .base_llm import BaseLLM -from ..enumeration import ChunkEnum -from ..schema import Message, StreamChunk, ToolCall - - -class OpenAILLM(BaseLLM): - """Asynchronous LLM client for OpenAI-compatible APIs supporting streaming completions and tool execution.""" - - def __init__(self, **kwargs): - """Initialize the OpenAI async client with API credentials and model configuration.""" - super().__init__(**kwargs) - - # Lazy client initialization - self._client = None - - def _create_client(self): - """Create and return an instance of the AsyncOpenAI client.""" - return AsyncOpenAI(api_key=self.api_key, base_url=self.base_url) - - @property - def client(self): - """Lazily create and return the AsyncOpenAI client.""" - if self._client is None: - self._client = self._create_client() - return self._client - - def _build_stream_kwargs( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - log_params: bool = True, - model_name: str | None = None, - **kwargs, - ) -> dict: - """Construct the parameter dictionary for the OpenAI Chat Completions API call. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - log_params: Whether to log parameters - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ - # Use the provided model_name or fall back to self.model_name - effective_model = model_name if model_name is not None else self.model_name - - # Construct the API parameters by merging multiple sources - llm_kwargs = { - "model": effective_model, - "messages": [x.simple_dump() for x in messages], - "tools": [x.simple_input_dump() for x in tools] if tools else None, - "stream": True, - **self.kwargs, - **kwargs, - } - - # Log parameters for debugging, with message/tool counts instead of full content - if log_params: - log_kwargs: dict = {} - for k, v in llm_kwargs.items(): - if k in ["messages", "tools"]: - log_kwargs[k] = len(v) if v is not None else 0 - else: - log_kwargs[k] = v - logger.info(f"llm_kwargs={log_kwargs}") - - return llm_kwargs - - async def _stream_chat( - self, - messages: list[Message], - tools: list[ToolCall] | None, - stream_kwargs: dict, - ) -> AsyncGenerator[StreamChunk, None]: - """Generate a stream of chat completion chunks including text, reasoning content, and tool calls.""" - stream_kwargs = stream_kwargs or {} - completion = await self.client.chat.completions.create(**stream_kwargs) - ret_tool_calls: list[ToolCall] = [] - - async for chunk in completion: - if not chunk.choices: - if hasattr(chunk, "usage") and chunk.usage: - yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump()) - continue - - delta = chunk.choices[0].delta - - if hasattr(delta, "reasoning_content") and delta.reasoning_content: - yield StreamChunk(chunk_type=ChunkEnum.THINK, chunk=delta.reasoning_content) - - if delta.content: - yield StreamChunk(chunk_type=ChunkEnum.ANSWER, chunk=delta.content) - - if delta.tool_calls is not None: - for tool_call in delta.tool_calls: - self._accumulate_tool_call_chunk(tool_call, ret_tool_calls) - - for tool_data in self._validate_and_serialize_tools(ret_tool_calls, tools): - yield StreamChunk(chunk_type=ChunkEnum.TOOL, chunk=tool_data) - - async def close(self): - """Asynchronously close the OpenAI client and release network resources.""" - if self._client is not None: - await self._client.close() - self._client = None - await super().close() diff --git a/reme/core/llm/openai_llm_sync.py b/reme/core/llm/openai_llm_sync.py deleted file mode 100644 index dadd1ca2..00000000 --- a/reme/core/llm/openai_llm_sync.py +++ /dev/null @@ -1,56 +0,0 @@ -"""Synchronous OpenAI-compatible LLM implementation supporting streaming, tool calls, and reasoning content.""" - -from typing import Generator - -from openai import OpenAI - -from .openai_llm import OpenAILLM -from ..enumeration import ChunkEnum -from ..schema import Message, StreamChunk, ToolCall - - -class OpenAILLMSync(OpenAILLM): - """Synchronous LLM client for OpenAI-compatible APIs, inheriting from OpenAILLM.""" - - def _create_client(self): - """Create and return an instance of the synchronous OpenAI client.""" - return OpenAI(api_key=self.api_key, base_url=self.base_url) - - def _stream_chat_sync( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - stream_kwargs: dict | None = None, - ) -> Generator[StreamChunk, None, None]: - """Synchronously generate a stream of chat completion chunks including text, reasoning, and tool calls.""" - stream_kwargs = stream_kwargs or {} - completion = self.client.chat.completions.create(**stream_kwargs) - ret_tool_calls: list[ToolCall] = [] - - for chunk in completion: - if not chunk.choices: - if hasattr(chunk, "usage") and chunk.usage: - yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump()) - continue - - delta = chunk.choices[0].delta - - if hasattr(delta, "reasoning_content") and delta.reasoning_content: - yield StreamChunk(chunk_type=ChunkEnum.THINK, chunk=delta.reasoning_content) - - if delta.content: - yield StreamChunk(chunk_type=ChunkEnum.ANSWER, chunk=delta.content) - - if delta.tool_calls is not None: - for tool_call in delta.tool_calls: - self._accumulate_tool_call_chunk(tool_call, ret_tool_calls) - - for tool_data in self._validate_and_serialize_tools(ret_tool_calls, tools): - yield StreamChunk(chunk_type=ChunkEnum.TOOL, chunk=tool_data) - - def close_sync(self): - """Close the synchronous OpenAI client and release network resources.""" - if self._client is not None: - self._client.close() - self._client = None - super().close_sync() diff --git a/reme/core/op/__init__.py b/reme/core/op/__init__.py deleted file mode 100644 index c99bed3a..00000000 --- a/reme/core/op/__init__.py +++ /dev/null @@ -1,24 +0,0 @@ -"""op""" - -from .base_op import BaseOp -from .base_ray_op import BaseRayOp -from .base_react import BaseReact -from .base_react_stream import BaseReactStream -from .base_tool import BaseTool -from .mcp_tool import MCPTool -from .parallel_op import ParallelOp -from .sequential_op import SequentialOp -from ..registry_factory import R - -__all__ = [ - "BaseOp", - "BaseRayOp", - "BaseReact", - "BaseReactStream", - "BaseTool", - "MCPTool", - "ParallelOp", - "SequentialOp", -] - -R.ops.register(MCPTool) diff --git a/reme/core/op/base_op.py b/reme/core/op/base_op.py deleted file mode 100644 index b016bb47..00000000 --- a/reme/core/op/base_op.py +++ /dev/null @@ -1,427 +0,0 @@ -"""Base operator class for LLM workflow execution and composition.""" - -import asyncio -import copy -import inspect -from abc import ABCMeta -from pathlib import Path -from typing import Callable, Optional, Any - -from agentscope.formatter import FormatterBase -from agentscope.model import ChatModelBase -from agentscope.token import HuggingFaceTokenCounter -from loguru import logger -from tqdm import tqdm - -from ..embedding import BaseEmbeddingModel -from ..file_store import BaseFileStore -from ..llm import BaseLLM -from ..prompt_handler import PromptHandler -from ..runtime_context import RuntimeContext -from ..schema import Response, ServiceConfig -from ..schema.service_config import OpConfig -from ..service_context import ServiceContext -from ..token_counter import BaseTokenCounter -from ..utils import camel_to_snake, CacheHandler, timer -from ..vector_store import BaseVectorStore - - -class BaseOp(metaclass=ABCMeta): - """Base operator class for LLM workflow execution and composition.""" - - __alias_name__: str = "" - - def __new__(cls, *args, **kwargs): - """Capture initialization arguments for object cloning.""" - instance = super().__new__(cls) - instance._init_args = copy.copy(args) - instance._init_kwargs = copy.copy(kwargs) - return instance - - def __init__( - self, - name: str = "", - async_mode: bool = True, - language: str = "", - prompt_name: str = "", - prompt_path: str = "", - as_llm: str | ChatModelBase = "default", - as_llm_formatter: str | FormatterBase = "default", - as_token_counter: str | HuggingFaceTokenCounter = "default", - llm: str | BaseLLM = "default", - embedding_model: str | BaseEmbeddingModel = "default", - vector_store: str | BaseVectorStore = "default", - file_store: str | BaseFileStore = "default", - token_counter: str | BaseTokenCounter = "default", - enable_cache: bool = False, - cache_path: str = "cache/op", - cache_expire_hours: float | None = None, - sub_ops: dict[str, "BaseOp"] | list["BaseOp"] | Optional["BaseOp"] = None, - input_mapping: dict[str, str] | None = None, - output_mapping: dict[str, str] | None = None, - enable_parallel: bool = False, - max_retries: int = 1, - raise_exception: bool = False, - **kwargs, - ): - """Initialize operator configurations and internal state.""" - self.name = name or self.__alias_name__ or camel_to_snake(self.__class__.__name__) - self.async_mode = async_mode - self.language = language - self.prompt = self._get_prompt_handler(prompt_name, prompt_path) - - self._as_llm = as_llm - self._as_llm_formatter = as_llm_formatter - self._as_token_counter = as_token_counter - self._llm = llm - self._embedding_model = embedding_model - self._vector_store = vector_store - self._file_store = file_store - self._token_counter = token_counter - - self.enable_cache = enable_cache - self.cache_path = cache_path - self.cache_expire_hours = cache_expire_hours - - self.sub_ops: list["BaseOp"] = [] - self.add_sub_ops(sub_ops) - - self.input_mapping = input_mapping - self.output_mapping = output_mapping - self.enable_parallel = enable_parallel # Control whether to execute tasks in parallel - self.max_retries = max(1, max_retries) - self.raise_exception = raise_exception - self.op_params = kwargs - - self._pending_tasks: list = [] - self.context: RuntimeContext | None = None - self._cache: CacheHandler | None = None - - def _get_prompt_handler(self, prompt_name: str, prompt_path: str) -> PromptHandler: - """Load prompt configuration from the associated YAML file.""" - if prompt_path: - path = Path(prompt_path) - else: - path = Path(inspect.getfile(self.__class__)) - if prompt_name: - path = path.with_stem(prompt_name) - return PromptHandler(language=self.language).load_prompt_by_file(path.with_suffix(".yaml")) - - def _handle_failure(self, e: Exception, attempt: int) -> str | None: - """Log failures and handle final retry logic.""" - message = f"[{self.__class__.__name__}] failed (attempt {attempt + 1}): {e}" - if attempt == self.max_retries - 1: - logger.exception(message) - if self.raise_exception: - raise e - return f"[{self.__class__.__name__}] failed: {e}" - else: - logger.warning(message) - return None - - @property - def cache(self) -> CacheHandler: - """Access the operator-specific cache handler.""" - assert self.enable_cache, "Cache is disabled!" - if not self._cache: - self._cache = CacheHandler(f"{self.cache_path}/{self.name}") - return self._cache - - @property - def service_context(self) -> ServiceContext: - """Access the service context.""" - assert self.context, "Service context is not initialized!" - return self.context.service_context - - @property - def service_config(self) -> ServiceConfig: - """Access the service configuration.""" - return self.service_context.service_config - - @property - def as_llm(self) -> ChatModelBase: - """Get the AgentScope LLM instance from ServiceContext.""" - if isinstance(self._as_llm, str): - self._as_llm = self.service_context.as_llms[self._as_llm] - return self._as_llm - - @property - def as_llm_formatter(self) -> FormatterBase: - """Get the AgentScope LLM formatter instance from ServiceContext.""" - if isinstance(self._as_llm_formatter, str): - self._as_llm_formatter = self.service_context.as_llm_formatters[self._as_llm_formatter] - return self._as_llm_formatter - - @property - def as_token_counter(self) -> HuggingFaceTokenCounter: - """Get the token counter instance from ServiceContext.""" - if isinstance(self._as_token_counter, str): - self._as_token_counter = self.service_context.as_token_counters[self._as_token_counter] - return self._as_token_counter - - @property - def llm(self) -> BaseLLM: - """Get the LLM instance from ServiceContext.""" - if isinstance(self._llm, str): - self._llm = self.service_context.llms[self._llm] - return self._llm - - @property - def embedding_model(self) -> BaseEmbeddingModel: - """Get the embedding model instance from ServiceContext.""" - if isinstance(self._embedding_model, str): - self._embedding_model = self.service_context.embedding_models[self._embedding_model] - return self._embedding_model - - @property - def vector_store(self) -> BaseVectorStore: - """Lazily initialize and return the vector store instance.""" - if isinstance(self._vector_store, str): - self._vector_store = self.service_context.vector_stores[self._vector_store] - return self._vector_store - - @property - def file_store(self) -> BaseFileStore: - """Lazily initialize and return the file store instance.""" - if isinstance(self._file_store, str): - self._file_store = self.service_context.file_stores[self._file_store] - return self._file_store - - @property - def token_counter(self) -> BaseTokenCounter: - """Get the token counter instance from ServiceContext.""" - if isinstance(self._token_counter, str): - self._token_counter = self.service_context.token_counters[self._token_counter] - return self._token_counter - - @property - def service_metadata(self) -> dict: - """Get service configuration metadata.""" - return self.service_context.service_config.metadata - - @property - def response(self) -> Response: - """Access the response object.""" - return self.context.response - - def before_execute_sync(self): - """Prepare context and validate before sync execution. - - This method performs the following steps: - 1. Apply input mapping to transform context variables - 2. Load operator-specific configuration from service config if available - 3. Override operator parameters and prompts based on config - """ - self.context.apply_mapping(self.input_mapping) - - if self.context.service_context is None: - return - - service_config = self.service_context.service_config - if self.name not in service_config.ops: - return - - op_config: OpConfig = service_config.ops[self.name] - - # Override operator parameters from config - if op_config.params: - for k, v in op_config.params.items(): - if hasattr(self, k): - setattr(self, k, v) - logger.info(f"[{self.__class__.__name__}] Set attribute '{k}' = {v}") - else: - self.op_params[k] = v - logger.info(f"[{self.__class__.__name__}] Set op_param '{k}' = {v}") - - # Load custom prompt templates from config - if op_config.prompt_dict: - self.prompt.load_prompt_dict(op_config.prompt_dict) - logger.info(f"[{self.__class__.__name__}] Loaded prompt keys={list(op_config.prompt_dict.keys())}") - - async def before_execute(self): - """Prepare context and validate before async execution.""" - self.before_execute_sync() - - def execute_sync(self): - """Define core sync logic in subclasses.""" - - async def execute(self): - """Define core async logic in subclasses.""" - - def after_execute_sync(self, response: Any): - """Finalize context and mappings after sync execution.""" - self.context.apply_mapping(self.output_mapping) - if response is not None: - if isinstance(response, dict): - for k, v in response.items(): - if k == "answer": - self.response.answer = v - elif k == "success": - self.response.success = v if isinstance(v, bool) else v.lower() == "true" - else: - self.response.metadata[k] = v - else: - self.response.answer = response - return response - - async def after_execute(self, output: Any): - """Finalize context and mappings after async execution.""" - return self.after_execute_sync(output) - - @timer - def call_sync(self, context: RuntimeContext = None, **kwargs): - """Execute the operator synchronously with retry logic.""" - self.context = RuntimeContext.from_context(context, **kwargs) - response = None - for i in range(self.max_retries): - try: - self.before_execute_sync() - response = self.execute_sync() - response = self.after_execute_sync(response) - break - except Exception as e: - response = self._handle_failure(e, i) - - return response - - @timer - async def call(self, context: RuntimeContext = None, **kwargs): - """Execute the operator asynchronously with retry logic.""" - self.context = RuntimeContext.from_context(context, **kwargs) - response = None - for i in range(self.max_retries): - try: - await self.before_execute() - response = await self.execute() - response = await self.after_execute(response) - break - except Exception as e: - response = self._handle_failure(e, i) - return response - - def submit_sync_task(self, fn: Callable, *args, **kwargs) -> "BaseOp": - """Submit a task to the thread pool or local queue.""" - if self.enable_parallel and self.service_context.thread_pool is not None: - task = self.service_context.thread_pool.submit(fn, *args, **kwargs) - else: - task = (fn, args, kwargs) - self._pending_tasks.append(task) - return self - - def submit_async_task(self, coro_fn: Callable, *args, **kwargs) -> "BaseOp": - """Submit an async task to the pending tasks queue.""" - task = coro_fn(*args, **kwargs) - self._pending_tasks.append(task) - return self - - def join_sync_tasks(self, task_desc: str = None) -> list: - """Wait for all pending sync tasks and return flattened results.""" - results = [] - for task in tqdm(self._pending_tasks, desc=task_desc or self.name): - if self.enable_parallel: - result = task.result() - else: - result = task[0](*task[1], **task[2]) - if result: - if isinstance(result, list): - results.extend(result) - else: - results.append(result) - self._pending_tasks.clear() - return results - - async def join_async_tasks(self, return_exceptions: bool = True) -> list: - """Wait for all pending async tasks and aggregate results.""" - if self.enable_parallel: - raw_results = await asyncio.gather(*self._pending_tasks, return_exceptions=return_exceptions) - else: - raw_results = [] - for task in self._pending_tasks: - try: - result = await task - raw_results.append(result) - except Exception as e: - if return_exceptions: - raw_results.append(e) - else: - raise - - results = [] - for result in raw_results: - if isinstance(result, Exception): - logger.error(f"[{self.__class__.__name__}] Async task failed: {result}") - elif result: - if isinstance(result, list): - results.extend(result) - else: - results.append(result) - self._pending_tasks.clear() - return results - - def add_sub_ops(self, sub_ops: dict[str, "BaseOp"] | list["BaseOp"] | Optional["BaseOp"]): - """Add child operators to this operator's sub_ops.""" - if not sub_ops: - return - - if isinstance(sub_ops, dict): - for name, op in sub_ops.items(): - assert self.async_mode == op.async_mode, "Async mode mismatch!" - op.name = name - if self.language: - op.language = self.language - self.sub_ops.append(op) - - elif isinstance(sub_ops, list): - for op in sub_ops: - assert self.async_mode == op.async_mode, "Async mode mismatch!" - if self.language: - op.language = self.language - self.sub_ops.append(op) - - else: - assert self.async_mode == sub_ops.async_mode, "Async mode mismatch!" - if self.language: - sub_ops.language = self.language - self.sub_ops.append(sub_ops) - - def add_sub_op(self, sub_op: "BaseOp"): - """Add a single child operator to this operator's sub_ops.""" - self.sub_ops.append(sub_op) - - def __lshift__(self, ops): - """Operator overload for adding sub-operators.""" - self.add_sub_ops(ops) - return self - - def __rshift__(self, op: "BaseOp"): - """Operator overload for sequential execution composition.""" - from .sequential_op import SequentialOp - - seq = SequentialOp(sub_ops=[self], async_mode=self.async_mode) - seq.add_sub_ops(op.sub_ops if isinstance(op, SequentialOp) else op) - return seq - - def __or__(self, op: "BaseOp"): - """Operator overload for parallel execution composition.""" - from .parallel_op import ParallelOp - - par = ParallelOp(sub_ops=[self], async_mode=self.async_mode) - par.add_sub_ops(op.sub_ops if isinstance(op, ParallelOp) else op) - return par - - def prompt_format(self, prompt_name: str, **kwargs) -> str: - """Format a prompt template with provided keyword arguments.""" - return self.prompt.prompt_format(prompt_name=prompt_name, **kwargs) - - def get_prompt(self, prompt_name: str) -> str: - """Get a prompt template by name.""" - return self.prompt.get_prompt(prompt_name=prompt_name) - - def copy(self, **kwargs): - """Create a copy of this operator with optional parameter overrides.""" - copy_op = self.__class__(*self._init_args, **{**self._init_kwargs, **kwargs}) - if self.sub_ops: - copy_op.sub_ops.clear() - for op in self.sub_ops: - copy_op.add_sub_op(op.copy()) - return copy_op diff --git a/reme/core/op/base_ray_op.py b/reme/core/op/base_ray_op.py deleted file mode 100644 index 2086d953..00000000 --- a/reme/core/op/base_ray_op.py +++ /dev/null @@ -1,124 +0,0 @@ -"""Base class for Ray-based parallel operations.""" - -from abc import ABCMeta -from typing import Callable - -import pandas as pd -from loguru import logger -from tqdm import tqdm - -from .base_op import BaseOp -from ..base_dict import BaseDict - -_RAY_IMPORT_ERROR: Exception | None = None - -try: - import ray -except Exception as _e: - _RAY_IMPORT_ERROR = _e - ray = None - - -class BaseRayOp(BaseOp, metaclass=ABCMeta): - """Base class for Ray-based parallel operations.""" - - def __init__(self, **kwargs): - if _RAY_IMPORT_ERROR: - raise ImportError("Ray requires extra dependencies. Install with `pip install ray`") - - super().__init__(**kwargs) - self._ray_task_list: list = [] - - def submit_and_join_parallel_op(self, op: BaseOp, **kwargs) -> list: - """Submit a BaseOp to be executed in parallel via Ray.""" - return self.submit_and_join_ray_task(fn=op.call, task_desc=op.name, context=self.context, **kwargs) - - def submit_and_join_ray_task(self, fn: Callable, parallel_key: str = "", task_desc: str = "", **kwargs) -> list: - """Divide data into chunks and execute them across Ray workers.""" - max_workers = self.service_context.ray_max_workers - self._ray_task_list.clear() - - # Automatically detect the key containing the list to parallelize - if not parallel_key: - for key, value in kwargs.items(): - if isinstance(value, list): - parallel_key = key - break - - if not parallel_key: - raise ValueError("No list found in kwargs to parallelize over.") - - parallel_list = kwargs.pop(parallel_key) - logger.info(f"Parallelizing '{parallel_key}' across {max_workers} workers") - - # Put large shared objects into the Ray Object Store once - optimized_kwargs = { - k: (ray.put(v) if isinstance(v, (pd.DataFrame, pd.Series, dict, list, BaseDict)) else v) - for k, v in kwargs.items() - } - - # Submit sliced chunks to reduce internode data transfer - remote_task_loop = ray.remote(self._ray_task_loop) - for i in range(max_workers): - chunk = parallel_list[i::max_workers] - if not chunk: - continue - - task = remote_task_loop.remote( - fn, - parallel_key, - chunk, - i, - **optimized_kwargs, - ) - self._ray_task_list.append(task) - logger.info(f"Submitted task {i + 1}/{max_workers} for {task_desc}") - - return self.join_ray_task(task_desc=task_desc) - - @staticmethod - def _ray_task_loop(internal_fn: Callable, parallel_key: str, chunk: list, actor_index: int, **kwargs) -> list: - """Execute the function over a specific chunk of data on a worker.""" - results = [] - for value in chunk: - current_kwargs = {**kwargs, "actor_index": actor_index, parallel_key: value} - t_result = internal_fn(**current_kwargs) - - if t_result is not None: - if isinstance(t_result, list): - results.extend(t_result) - else: - results.append(t_result) - return results - - def submit_ray_task(self, fn, *args, **kwargs): - """Submit a single Ray task to the task list for later execution.""" - if not ray.is_initialized(): - ray.init(num_cpus=self.service_context.ray_max_workers, ignore_reinit_error=True) - - remote_fn = ray.remote(fn) - task = remote_fn.remote(*args, **kwargs) - self._ray_task_list.append(task) - return self - - def join_ray_task(self, task_desc: str | None = None) -> list: - """Collect results from Ray workers using a progress bar.""" - results = [] - unfinished = list(self._ray_task_list) - - with tqdm(total=len(unfinished), desc=task_desc or f"{self.name}_ray") as pbar: - while unfinished: - ready, unfinished = ray.wait(unfinished, num_returns=1) - for obj_ref in ready: - try: - t_result = ray.get(obj_ref) - if isinstance(t_result, list): - results.extend(t_result) - elif t_result is not None: - results.append(t_result) - except Exception as e: - logger.error(f"Worker task failed: {e}") - pbar.update(1) - - self._ray_task_list.clear() - return results diff --git a/reme/core/op/base_react.py b/reme/core/op/base_react.py deleted file mode 100644 index c6587204..00000000 --- a/reme/core/op/base_react.py +++ /dev/null @@ -1,176 +0,0 @@ -"""Base memory agent for handling memory operations with tool-based reasoning.""" - -import asyncio -from typing import TYPE_CHECKING - -from loguru import logger - -from ..enumeration import Role -from ..op import BaseOp -from ..schema import Message - -if TYPE_CHECKING: - from . import BaseTool - - -class BaseReact(BaseOp): - """ReAct agent that performs reasoning and acting cycles with tools.""" - - def __init__( - self, - tools: list["BaseTool"], - tool_call_interval: float = 0, - max_steps: int = 10, - **kwargs, - ): - """Initialize ReAct agent with tools and execution parameters.""" - kwargs["sub_ops"] = tools or [] - super().__init__(**kwargs) - # Filter only BaseTool instances from sub_ops - from . import BaseTool - - self.sub_ops: list[BaseTool] = [t for t in self.sub_ops if isinstance(t, BaseTool)] - self.tool_call_interval: float = tool_call_interval - self.max_steps: int = max_steps - - @property - def tools(self) -> list["BaseTool"]: - """Return available tools for the agent.""" - return self.sub_ops - - def pop_tool(self, name: str) -> "BaseTool | None": - """Remove and return a tool from self.tools by name.""" - for i, tool in enumerate(self.sub_ops): - if tool.tool_call.name == name: - return self.sub_ops.pop(i) - return None - - async def build_messages(self) -> list[Message]: - """Build initial message list from context query or messages.""" - if self.context.get("query"): - messages = [Message(role=Role.USER, content=self.context.query)] - elif self.context.get("messages"): - messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] - else: - raise ValueError("input must have either `query` or `messages`") - return messages - - async def _reasoning_step( - self, - messages: list[Message], - tools: list["BaseTool"], - step: int, - stage: str = "", - **kwargs, - ) -> tuple[Message, bool]: - """Execute one reasoning step where LLM decides whether to use tools.""" - # Get tool definitions for LLM - tool_calls = [t.tool_call for t in tools] - - # Generate assistant response with potential tool calls - assistant_message: Message = await self.llm.chat(messages=messages, tools=tool_calls, **kwargs) - messages.append(assistant_message) - assistant_content: str = assistant_message.simple_dump(as_dict=False) - logger.info(f"[{self.__class__.__name__} {stage or ''} step{step}] assistant={assistant_content}") - - # Determine if tools should be called - should_act = bool(assistant_message.tool_calls) - return assistant_message, should_act - - async def _acting_step( - self, - assistant_message: Message, - tools: list["BaseTool"], - step: int, - stage: str = "", - **kwargs, - ) -> tuple[list["BaseTool"], list[Message]]: - """Execute tool calls requested by the assistant and collect results.""" - tool_list: list["BaseTool"] = [] - tool_messages: list[Message] = [] - - if not assistant_message.tool_calls: - return tool_list, tool_messages - - # Create tool name to tool instance mapping - tool_dict = {t.tool_call.name: t for t in tools} - for j, tool_call in enumerate(assistant_message.tool_calls): - prefix: str = f"[{self.__class__.__name__} {stage or ''} step{step}.{j}]" - if tool_call.name not in tool_dict: - logger.warning(f"{prefix} unknown tool_call={tool_call.name}") - continue - - logger.info(f"{prefix} submit tool_call[{tool_call.name}] arguments={tool_call.arguments}") - - # Create independent tool copy with unique ID - tool_copy: BaseTool = tool_dict[tool_call.name].copy() - tool_copy.tool_call.id = tool_call.id - tool_list.append(tool_copy) - - # Create isolated kwargs for each tool call to avoid parameter conflicts - tool_kwargs = {**kwargs, **tool_call.argument_dict} - self.submit_async_task(tool_copy.call, service_context=self.service_context, **tool_kwargs) - if self.tool_call_interval > 0: - await asyncio.sleep(self.tool_call_interval) - - # Wait for all tool executions to complete - await self.join_async_tasks() - - # Collect tool results as messages - for j, tool in enumerate(tool_list): - tool_messages.append( - Message( - role=Role.TOOL, - content=tool.response.answer, - tool_call_id=tool.tool_call.id, - ), - ) - prefix: str = f"[{self.__class__.__name__} {stage or ''} step{step}.{j}]" - logger.info(f"{prefix} join tool={tool.name} result={tool.response.answer}") - return tool_list, tool_messages - - async def react(self, messages: list[Message], tools: list["BaseTool"], stage: str = ""): - """Run ReAct loop alternating between reasoning and acting until completion.""" - success: bool = False - used_tools: list[BaseTool] = [] - for step in range(self.max_steps): - # Reasoning: LLM decides next action - assistant_message, should_act = await self._reasoning_step(messages, tools, step=step, stage=stage) - - if not should_act: - # No tools requested, task complete - success = True - break - - # Acting: execute tools and collect results - t_tools, tool_messages = await self._acting_step(assistant_message, tools, step=step, stage=stage) - used_tools.extend(t_tools) - messages.extend(tool_messages) - - return used_tools, messages, success - - async def execute(self): - """Execute the ReAct agent and return final results.""" - # Log available tools - for i, tool in enumerate(self.tools): - logger.info(f"[{self.__class__.__name__}] {i}.tool_call={tool.tool_call.simple_input_dump(as_dict=False)}") - - # Build and log initial messages - messages = await self.build_messages() - for i, message in enumerate(messages): - role = message.name or message.role - logger.info(f"[{self.__class__.__name__}] role={role} {message.simple_dump(as_dict=False)}") - - # Run ReAct loop - t_tools, messages, success = await self.react(messages, self.tools) - - # Get the last assistant message as the final answer - assistant_messages = [m for m in messages if m.role == Role.ASSISTANT] - answer = assistant_messages[-1].content if assistant_messages else "" - - return { - "answer": answer, - "success": success, - "messages": messages, - "tools": t_tools, - } diff --git a/reme/core/op/base_react_stream.py b/reme/core/op/base_react_stream.py deleted file mode 100644 index c148d574..00000000 --- a/reme/core/op/base_react_stream.py +++ /dev/null @@ -1,249 +0,0 @@ -"""Base memory agent for handling memory operations with tool-based reasoning.""" - -import asyncio -from typing import TYPE_CHECKING - -from loguru import logger - -from ..enumeration import Role, ChunkEnum -from ..op import BaseOp -from ..schema import Message, StreamChunk - -if TYPE_CHECKING: - from . import BaseTool - - -class BaseReactStream(BaseOp): - """ReAct agent that performs reasoning and acting cycles with tools.""" - - def __init__( - self, - tools: list["BaseTool"], - tool_call_interval: float = 0, - max_steps: int = 10, - **kwargs, - ): - """Initialize ReAct agent with tools and execution parameters.""" - kwargs["sub_ops"] = tools or [] - super().__init__(**kwargs) - # Filter only BaseTool instances from sub_ops - from . import BaseTool - - self.sub_ops: list[BaseTool] = [t for t in self.sub_ops if isinstance(t, BaseTool)] - self.tool_call_interval: float = tool_call_interval - self.max_steps: int = max_steps - - @property - def tools(self) -> list["BaseTool"]: - """Return available tools for the agent.""" - return self.sub_ops - - def pop_tool(self, name: str) -> "BaseTool | None": - """Remove and return a tool from self.tools by name.""" - for i, tool in enumerate(self.sub_ops): - if tool.tool_call.name == name: - return self.sub_ops.pop(i) - return None - - async def build_messages(self) -> list[Message]: - """Build initial message list from context query or messages.""" - if self.context.get("query"): - messages = [Message(role=Role.USER, content=self.context.query)] - elif self.context.get("messages"): - messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] - else: - raise ValueError("input must have either `query` or `messages`") - return messages - - async def _reasoning_step( - self, - messages: list[Message], - tools: list["BaseTool"], - step: int, - stage: str = "", - **kwargs, - ) -> tuple[Message, bool]: - """Execute one reasoning step where LLM decides whether to use tools.""" - tool_calls = [t.tool_call for t in tools] - - start_chunk = StreamChunk(chunk_type=ChunkEnum.STEP_START, metadata={"step": step, "stage": stage}) - await self.context.add_stream_chunk(start_chunk) - - # State for accumulating message content from stream - state = { - "reasoning_content": "", - "content": "", - "tool_calls": [], - } - - async for stream_chunk in self.llm.stream_chat(messages=messages, tools=tool_calls, **kwargs): # noqa - if stream_chunk.chunk_type in [ChunkEnum.ANSWER, ChunkEnum.THINK, ChunkEnum.ERROR]: - await self.context.add_stream_chunk(stream_chunk) - - # Accumulate content based on chunk type - if stream_chunk.chunk_type is ChunkEnum.THINK: - state["reasoning_content"] += stream_chunk.chunk - - elif stream_chunk.chunk_type is ChunkEnum.ANSWER: - state["content"] += stream_chunk.chunk - - elif stream_chunk.chunk_type is ChunkEnum.TOOL: - state["tool_calls"].append(stream_chunk.chunk) - - # Build the final assistant message from accumulated state - assistant_message = Message(role=Role.ASSISTANT, **state) - messages.append(assistant_message) - logger.info( - f"[{self.__class__.__name__} {stage or ''} step{step}] " - f"assistant={assistant_message.simple_dump(as_dict=False)}", - ) - - should_act = bool(assistant_message.tool_calls) - return assistant_message, should_act - - async def _acting_step( - self, - assistant_message: Message, - tools: list["BaseTool"], - step: int, - stage: str = "", - **kwargs, - ) -> tuple[list["BaseTool"], list[Message]]: - """Execute tool calls serially and collect results with streaming output.""" - tool_list: list["BaseTool"] = [] - tool_messages: list[Message] = [] - - if not assistant_message.tool_calls: - return tool_list, tool_messages - - # Create tool name to tool instance mapping - tool_dict = {t.tool_call.name: t for t in tools} - - # Execute tools serially for better streaming experience - for j, tool_call in enumerate(assistant_message.tool_calls): - prefix: str = f"[{self.__class__.__name__} {stage or ''} step{step}.{j}]" - if tool_call.name not in tool_dict: - logger.warning(f"{prefix} unknown tool_call={tool_call.name}") - # Emit error chunk for unknown tool - await self.context.add_stream_chunk( - StreamChunk( - chunk_type=ChunkEnum.ERROR, - chunk=f"Unknown tool: {tool_call.name}", - metadata={"step": step, "tool_index": j, "tool_name": tool_call.name}, - ), - ) - continue - - logger.info(f"{prefix} submit tool_call[{tool_call.name}] arguments={tool_call.arguments}") - - # Emit tool execution start signal - await self.context.add_stream_chunk( - StreamChunk( - chunk_type=ChunkEnum.TOOL, - chunk=f"Executing tool: {tool_call.name} {tool_call.arguments}", - metadata={ - "step": step, - "tool_index": j, - "tool_name": tool_call.name, - "arguments": tool_call.arguments, - }, - ), - ) - - # Create independent tool copy with unique ID - tool_copy: BaseTool = tool_dict[tool_call.name].copy() - tool_copy.tool_call.id = tool_call.id - tool_list.append(tool_copy) - - # Create isolated kwargs for each tool call to avoid parameter conflicts - tool_kwargs = {**kwargs, **tool_call.argument_dict} - - # Execute tool serially (wait for completion before next tool) - await tool_copy.call(service_context=self.service_context, **tool_kwargs) - - # Get tool result immediately after execution - tool_result = tool_copy.response.answer - tool_messages.append( - Message( - role=Role.TOOL, - content=tool_result, - tool_call_id=tool_copy.tool_call.id, - ), - ) - logger.info(f"{prefix} tool={tool_copy.name} result={tool_result}") - - await self.context.add_stream_chunk( - StreamChunk( - chunk_type=ChunkEnum.TOOL_RESULT, - chunk=tool_result, - metadata={ - "step": step, - "tool_index": j, - "tool_name": tool_copy.name, - "tool_call_id": tool_copy.tool_call.id, - }, - ), - ) - - # Optional interval between tool calls - if self.tool_call_interval > 0 and j < len(assistant_message.tool_calls) - 1: - await asyncio.sleep(self.tool_call_interval) - - return tool_list, tool_messages - - async def react(self, messages: list[Message], tools: list["BaseTool"], stage: str = ""): - """Run ReAct loop alternating between reasoning and acting until completion.""" - success: bool = False - used_tools: list[BaseTool] = [] - for step in range(self.max_steps): - # Reasoning: LLM decides next action - assistant_message, should_act = await self._reasoning_step(messages, tools, step=step, stage=stage) - - if not should_act: - # No tools requested, task complete - success = True - break - - # Acting: execute tools and collect results - t_tools, tool_messages = await self._acting_step(assistant_message, tools, step=step, stage=stage) - used_tools.extend(t_tools) - messages.extend(tool_messages) - - return used_tools, messages, success - - async def execute(self): - """Execute the ReAct agent with streaming output and return final results.""" - for i, tool in enumerate(self.tools): - logger.info(f"[{self.__class__.__name__}] {i}.tool_call={tool.tool_call.simple_input_dump(as_dict=False)}") - - # Build and log initial messages - messages = await self.build_messages() - for i, message in enumerate(messages): - role = message.name or message.role - logger.info(f"[{self.__class__.__name__}] role={role} {message.simple_dump(as_dict=False)}") - - # Run ReAct loop with streaming - t_tools, messages, success = await self.react(messages, self.tools) - - # Emit final done signal - await self.context.add_stream_chunk( - StreamChunk( - chunk_type=ChunkEnum.DONE, - chunk="", - metadata={ - "success": success, - "total_steps": len(t_tools), - }, - ), - ) - - # Get the last assistant message as the final answer - assistant_messages = [m for m in messages if m.role == Role.ASSISTANT] - answer = assistant_messages[-1].content if assistant_messages else "" - - return { - "answer": answer, - "success": success, - "messages": messages, - "tools": t_tools, - } diff --git a/reme/core/op/base_tool.py b/reme/core/op/base_tool.py deleted file mode 100644 index 5b41b1fb..00000000 --- a/reme/core/op/base_tool.py +++ /dev/null @@ -1,58 +0,0 @@ -"""Base class for tools""" - -from abc import ABCMeta - -from . import BaseOp -from ..schema import ToolCall - - -class BaseTool(BaseOp, metaclass=ABCMeta): - """Base class for tools""" - - def __init__(self, **kwargs): - super().__init__(**kwargs) - self._tool_call: ToolCall | None = None - - def _build_tool_call(self) -> ToolCall: - """Build and return the tool call schema; override in subclasses.""" - - def _validate_inputs(self): - """Validate the inputs.""" - parameters = self.tool_call.parameters - if parameters.type == "object" and parameters.properties: - required_list = parameters.required or [] - required_keys = {k: (k in required_list) for k in parameters.properties.keys()} - self.context.validate_required_keys(required_keys, self.name) - - @property - def tool_call(self) -> ToolCall: - """Get the tool call schema.""" - if self._tool_call is None: - self._tool_call = self._build_tool_call() - self._tool_call.name = self._tool_call.name or self.name - return self._tool_call - - def set_tool_call(self, tool_call: ToolCall | dict): - """Set the tool call schema.""" - if isinstance(tool_call, dict): - self._tool_call = ToolCall(**tool_call) - elif isinstance(tool_call, ToolCall): - self._tool_call = tool_call - else: - raise ValueError(f"Invalid tool call: {tool_call}") - - self._tool_call.name = self._tool_call.name or self.name - - @property - def input_dict(self) -> dict: - """Get the input dict.""" - parameters = self.tool_call.parameters - if parameters.type != "object" or not parameters.properties: - return {} - required_keys = set(parameters.required or []) - return {k: self.context[k] for k in parameters.properties.keys() if (k in required_keys or k in self.context)} - - def before_execute_sync(self): - """Hook before execute""" - super().before_execute_sync() - self._validate_inputs() diff --git a/reme/core/op/mcp_tool.py b/reme/core/op/mcp_tool.py deleted file mode 100644 index bcd64ad2..00000000 --- a/reme/core/op/mcp_tool.py +++ /dev/null @@ -1,83 +0,0 @@ -"""MCP (Model Context Protocol) tool integration for remote tool execution.""" - -from mcp.types import CallToolResult, TextContent - -from .base_tool import BaseTool -from ..schema import ToolCall -from ..utils import MCPClient - - -class MCPTool(BaseTool): - """Operator for calling remote MCP (Model Context Protocol) tools.""" - - def __init__( - self, - mcp_server: str = "", - tool_name: str = "", - parameter_required: list[str] | None = None, - parameter_optional: list[str] | None = None, - parameter_deleted: list[str] | None = None, - max_retries: int = 3, - timeout: float | None = None, - raise_exception: bool = False, - **kwargs, - ): - super().__init__(max_retries=max_retries, raise_exception=raise_exception, **kwargs) - - self.mcp_server: str = mcp_server - self.tool_name: str = tool_name - self.parameter_required: list[str] | None = parameter_required - self.parameter_optional: list[str] | None = parameter_optional - self.parameter_deleted: list[str] | None = parameter_deleted - self.timeout: float | None = timeout - - # Example MCP marketplace: https://bailian.console.aliyun.com/?tab=mcp#/mcp-market - self._client: MCPClient | None = None - - @property - def client(self) -> MCPClient: - """Lazily initialize and return the MCP client.""" - if self._client is None: - self._client = MCPClient(self.service_context.service_config.mcp_servers) - return self._client - - def _build_tool_call(self) -> ToolCall: - tool_call_dict = self.service_context.mcp_server_mapping[self.mcp_server] - tool_call: ToolCall = tool_call_dict[self.tool_name].model_copy(deep=True) - - # Initialize required list if not exists - if tool_call.parameters.required is None: - tool_call.parameters.required = [] - - if self.parameter_required: - for name in self.parameter_required: - if name not in tool_call.parameters.required: - tool_call.parameters.required.append(name) - - if self.parameter_optional: - for name in self.parameter_optional: - if name in tool_call.parameters.required: - tool_call.parameters.required.remove(name) - - if self.parameter_deleted: - for name in self.parameter_deleted: - tool_call.parameters.properties.pop(name, None) - if tool_call.parameters.required and name in tool_call.parameters.required: - tool_call.parameters.required.remove(name) - - return tool_call - - async def execute(self): - tool_result: CallToolResult = await self.client.call_tool( - server_name=self.mcp_server, - tool_name=self.tool_name, - arguments=self.input_dict, - ) - self.context.tool_result = tool_result - - text_result = [] - for block in tool_result.content: - if isinstance(block, TextContent): - text_result.append(block.text) - output: str = "\n".join(text_result) - return output diff --git a/reme/core/op/parallel_op.py b/reme/core/op/parallel_op.py deleted file mode 100644 index 8bca5790..00000000 --- a/reme/core/op/parallel_op.py +++ /dev/null @@ -1,33 +0,0 @@ -"""Module providing the ParallelOp class for concurrent operation execution.""" - -from .base_op import BaseOp - - -class ParallelOp(BaseOp): - """Operation class that executes multiple sub-operations in parallel.""" - - async def execute(self): - """Executes all sub-operations concurrently using asynchronous tasks.""" - for op in self.sub_ops: - assert op.async_mode - self.submit_async_task(op.call, context=self.context) - return await self.join_async_tasks() - - def execute_sync(self): - """Executes all sub-operations concurrently using synchronous task management.""" - for op in self.sub_ops: - assert not op.async_mode - self.submit_sync_task(op.call_sync, context=self.context) - return self.join_sync_tasks() - - def __lshift__(self, op: dict[str, BaseOp] | list[BaseOp] | BaseOp): - """Raises RuntimeError as the shift operator is not supported for parallel operations.""" - raise RuntimeError(f"`<<` is not supported in `{self.name}`") - - def __or__(self, op: BaseOp): - """Adds sub-operations to the current parallel group using the bitwise OR operator.""" - if isinstance(op, ParallelOp) and op.sub_ops: - self.add_sub_ops(op.sub_ops) - else: - self.add_sub_op(op) - return self diff --git a/reme/core/op/sequential_op.py b/reme/core/op/sequential_op.py deleted file mode 100644 index fed44dc3..00000000 --- a/reme/core/op/sequential_op.py +++ /dev/null @@ -1,35 +0,0 @@ -"""Module providing the SequentialOp class for serial operation execution.""" - -from .base_op import BaseOp - - -class SequentialOp(BaseOp): - """Operation class that executes sub-operations one after another in order.""" - - async def execute(self): - """Executes sub-operations sequentially using asynchronous awaits.""" - result = None - for op in self.sub_ops: - assert op.async_mode - result = await op.call(context=self.context) - return result - - def execute_sync(self): - """Executes sub-operations sequentially in a synchronous blocking manner.""" - result = None - for op in self.sub_ops: - assert not op.async_mode - result = op.call_sync(context=self.context) - return result - - def __lshift__(self, op: dict[str, BaseOp] | list[BaseOp] | BaseOp): - """Raises RuntimeError as the left shift operator is not supported.""" - raise RuntimeError(f"`<<` is not supported in `{self.name}`") - - def __rshift__(self, op: BaseOp): - """Appends operations to the sequence using the bitwise right shift operator.""" - if isinstance(op, SequentialOp) and op.sub_ops: - self.add_sub_ops(op.sub_ops) - else: - self.add_sub_op(op) - return self diff --git a/reme/core/prompt_handler.py b/reme/core/prompt_handler.py deleted file mode 100644 index 22e1b095..00000000 --- a/reme/core/prompt_handler.py +++ /dev/null @@ -1,146 +0,0 @@ -"""Module for managing and formatting prompt templates from files or dictionaries.""" - -import json -from pathlib import Path -from string import Formatter -from typing import Any, Dict, Optional, Union - -import yaml -from loguru import logger - -from .base_dict import BaseDict - - -class PromptHandler(BaseDict): - """A context-aware handler for loading, retrieving, and formatting prompt templates.""" - - def __init__(self, language: str = "", **kwargs): - super().__init__(**kwargs) - # Use object.__setattr__ to avoid storing 'language' in the dict - object.__setattr__(self, "language", language.strip()) - - def load_prompt_by_file( - self, - prompt_file_path: Optional[Union[Path, str]] = None, - overwrite: bool = True, - ) -> "PromptHandler": - """Load prompt configurations from a YAML or JSON file.""" - if prompt_file_path is None: - return self - - if isinstance(prompt_file_path, str): - prompt_file_path = Path(prompt_file_path) - - if not prompt_file_path.exists(): - return self - - suffix = prompt_file_path.suffix.lower() - - with prompt_file_path.open(encoding="utf-8") as f: - if suffix in [".yaml", ".yml"]: - prompt_dict = yaml.safe_load(f) - elif suffix == ".json": - prompt_dict = json.load(f) - else: - raise ValueError(f"Unsupported file format: {suffix}") - - self.load_prompt_dict(prompt_dict, overwrite=overwrite) - return self - - def load_prompt_dict( - self, - prompt_dict: Optional[Dict[str, Any]] = None, - overwrite: bool = True, - ) -> "PromptHandler": - """Merge a dictionary of prompt strings into the current context.""" - if not prompt_dict: - return self - - for key, value in prompt_dict.items(): - if not isinstance(value, str): - continue - if key in self: - if overwrite: - logger.warning(f"Overwriting prompt '{key}'") - self[key] = value - else: - self[key] = value - - return self - - def get_prompt(self, prompt_name: str, fallback_to_base: bool = True) -> str: - """Retrieve a prompt by name with automatic language suffix handling.""" - if self.language and not prompt_name.endswith(f"_{self.language}"): - key_with_lang = f"{prompt_name}_{self.language}" - if key_with_lang in self: - return self[key_with_lang].strip() - - if prompt_name in self: - return self[prompt_name].strip() - - if fallback_to_base and self.language and prompt_name.endswith(f"_{self.language}"): - base_name = prompt_name[: -(len(self.language) + 1)] - if base_name in self: - return self[base_name].strip() - - raise KeyError(f"Prompt '{prompt_name}' not found. Available: {list(self.keys())[:10]}") - - def has_prompt(self, prompt_name: str) -> bool: - """Check if a prompt exists.""" - try: - self.get_prompt(prompt_name) - return True - except KeyError: - return False - - def list_prompts(self, language_filter: Optional[str] = None) -> list[str]: - """List all available prompt names.""" - if language_filter is None: - return list(self.keys()) - suffix = f"_{language_filter.strip()}" - return [key for key in self.keys() if key.endswith(suffix)] - - @staticmethod - def _extract_format_fields(template: str) -> set[str]: - """Extract all format field names from a template string.""" - return {field_name for _, field_name, _, _ in Formatter().parse(template) if field_name is not None} - - @staticmethod - def _filter_conditional_lines(prompt: str, flags: Dict[str, bool]) -> str: - """Filter lines based on boolean flags.""" - filtered_lines = [] - for line in prompt.split("\n"): - matched_flag = None - for flag_name in flags: - if line.startswith(f"[{flag_name}]"): - matched_flag = flag_name - break - if matched_flag is None: - filtered_lines.append(line) - elif flags[matched_flag]: - filtered_lines.append(line[len(f"[{matched_flag}]") :]) - return "\n".join(filtered_lines) - - def prompt_format(self, prompt_name: str, validate: bool = True, **kwargs) -> str: - """Format a prompt with conditional line filtering and variable substitution.""" - prompt = self.get_prompt(prompt_name) - - flag_kwargs = {k: v for k, v in kwargs.items() if isinstance(v, bool)} - format_kwargs = {k: v for k, v in kwargs.items() if not isinstance(v, bool)} - - if flag_kwargs: - prompt = self._filter_conditional_lines(prompt, flag_kwargs) - - if validate: - required_fields = self._extract_format_fields(prompt) - missing_fields = required_fields - set(format_kwargs.keys()) - if missing_fields: - raise ValueError(f"Missing format variables for '{prompt_name}': {sorted(missing_fields)}") - - if format_kwargs: - prompt = prompt.format(**format_kwargs) - - return prompt.strip() - - def __repr__(self) -> str: - return f"PromptHandler(language='{self.language}', num_prompts={len(self)})" diff --git a/reme/core/registry_factory.py b/reme/core/registry_factory.py deleted file mode 100644 index 921f27a8..00000000 --- a/reme/core/registry_factory.py +++ /dev/null @@ -1,50 +0,0 @@ -"""Module providing a registry class for managing class-to-name mappings via decorators.""" - -import inspect -from typing import Callable, TypeVar - -from .base_dict import BaseDict -from .utils import singleton - -T = TypeVar("T") - - -class Registry(BaseDict): - """A registry container that uses decorators to map and store class references.""" - - def register(self, name: str | type = "") -> Callable[[type[T]], type[T]] | type[T]: - """Return a decorator that registers a class under a specific name in the registry.""" - if inspect.isclass(name): - self[name.__name__] = name - return name - - else: - - def decorator(cls): - key: str = name if isinstance(name, str) and name else cls.__name__ - self[key] = cls - return cls - - return decorator - - -@singleton -class RegistryFactory: - """A factory class for creating registries.""" - - def __init__(self): - self.llms = Registry() - self.as_llms = Registry() - self.as_llm_formatters = Registry() - self.as_token_counters = Registry() - self.embedding_models = Registry() - self.vector_stores = Registry() - self.file_stores = Registry() - self.ops = Registry() - self.flows = Registry() - self.services = Registry() - self.token_counters = Registry() - self.file_watchers = Registry() - - -R = RegistryFactory() diff --git a/reme/core/runtime_context.py b/reme/core/runtime_context.py deleted file mode 100644 index 43c7c7d8..00000000 --- a/reme/core/runtime_context.py +++ /dev/null @@ -1,87 +0,0 @@ -"""Runtime context for managing response states and asynchronous data streaming.""" - -import asyncio - -from .base_dict import BaseDict -from .enumeration import ChunkEnum -from .schema import Response, StreamChunk -from .service_context import ServiceContext - - -class RuntimeContext(BaseDict): - """Context for execution state, response metadata, and stream queues.""" - - def __init__( - self, - response: Response | None = None, - stream_queue: asyncio.Queue | None = None, - service_context: ServiceContext | None = None, - **kwargs, - ): - """Initialize the context with optional response and queue.""" - super().__init__(**kwargs) - self.response: Response | None = response or Response() - self.stream_queue: asyncio.Queue | None = stream_queue - self.service_context: ServiceContext | None = service_context - - @classmethod - def from_context(cls, context: "RuntimeContext | None" = None, **kwargs) -> "RuntimeContext": - """Create a new context from an existing instance or keywords.""" - if context is None: - return cls(**kwargs) - else: - if kwargs: - context.update(kwargs) - return context - - async def _enqueue(self, chunk: StreamChunk) -> None: - """Internal helper to put a chunk into the queue if it exists.""" - if self.stream_queue: - await self.stream_queue.put(chunk) - - async def add_stream_string(self, chunk: str, chunk_type: ChunkEnum) -> "RuntimeContext": - """Enqueue a stream chunk from a raw string and type.""" - await self._enqueue(StreamChunk(chunk_type=chunk_type, chunk=chunk)) - return self - - async def add_stream_chunk(self, stream_chunk: StreamChunk) -> "RuntimeContext": - """Enqueue an existing stream chunk.""" - await self._enqueue(stream_chunk) - return self - - async def add_stream_done(self) -> "RuntimeContext": - """Enqueue a termination chunk to signal the end of the stream.""" - await self._enqueue(StreamChunk(chunk_type=ChunkEnum.DONE, chunk="", done=True)) - return self - - def add_response_error(self, e: Exception) -> "RuntimeContext": - """Record an exception into the response object.""" - self.response.success = False - self.response.answer = str(e) - return self - - def apply_mapping(self, mapping: dict[str, str]) -> "RuntimeContext": - """Copy internal values based on a source-to-target key map.""" - if not mapping: - return self - - for source, target in mapping.items(): - if source in self: - self[target] = self[source] - return self - - def validate_required_keys( - self, - required_keys: dict[str, bool], - context_name: str = "context", - ) -> "RuntimeContext": - """Ensure all required keys are present in the context. - - Args: - required_keys: Dictionary mapping key names to boolean indicating if required - context_name: Name of the context for error messages (e.g., operator name) - """ - for key, is_required in required_keys.items(): - if is_required and key not in self: - raise ValueError(f"{context_name}: missing required input '{key}'") - return self diff --git a/reme/core/schema/__init__.py b/reme/core/schema/__init__.py deleted file mode 100644 index b6445a28..00000000 --- a/reme/core/schema/__init__.py +++ /dev/null @@ -1,59 +0,0 @@ -"""schema""" - -from .as_msg_stat import AsBlockStat, AsMsgStat -from .cut_point_result import CutPointResult -from .file_metadata import FileMetadata -from .memory_chunk import MemoryChunk -from .memory_node import MemoryNode -from .memory_search_result import MemorySearchResult -from .message import ContentBlock, Message, Trajectory -from .request import Request -from .response import Response -from .service_config import ( - CmdConfig, - EmbeddingModelConfig, - FileWatcherConfig, - FlowConfig, - HttpConfig, - LLMConfig, - MCPConfig, - FileStoreConfig, - ServiceConfig, - TokenCounterConfig, - VectorStoreConfig, -) -from .stream_chunk import StreamChunk -from .tool_call import ToolAttr, ToolCall -from .truncation_result import TruncationResult -from .vector_node import VectorNode - -__all__ = [ - "AsBlockStat", - "AsMsgStat", - "CutPointResult", - "CmdConfig", - "ContentBlock", - "EmbeddingModelConfig", - "FileMetadata", - "FileWatcherConfig", - "FlowConfig", - "HttpConfig", - "LLMConfig", - "MCPConfig", - "MemoryChunk", - "MemoryNode", - "MemorySearchResult", - "FileStoreConfig", - "Message", - "Request", - "Response", - "ServiceConfig", - "StreamChunk", - "TokenCounterConfig", - "Trajectory", - "ToolAttr", - "ToolCall", - "TruncationResult", - "VectorNode", - "VectorStoreConfig", -] diff --git a/reme/core/schema/as_msg_stat.py b/reme/core/schema/as_msg_stat.py deleted file mode 100644 index a8d2daaa..00000000 --- a/reme/core/schema/as_msg_stat.py +++ /dev/null @@ -1,96 +0,0 @@ -"""Schema definitions for AgentScope message statistics.""" - -from pydantic import BaseModel, Field - -_TRUNCATION_NOTICE_MARKER = "<<>>" - -_DEFAULT_MAX_BLOCK_TEXT_PREVIEW_LENGTH = 100 -_DEFAULT_MAX_FORMATTER_TEXT_LENGTH = 1000 - - -class AsBlockStat(BaseModel): - """Statistics and metadata for a single content block in an AgentScope message.""" - - block_type: str = Field(default=...) - text: str = Field(default="", description="Text content of the block") - token_count: int = Field(default=0, description="Token count of the block, including base64 data") - - # For tool_use and tool_result blocks - tool_name: str = Field(default="", description="Tool name for tool_use/tool_result blocks") - tool_input: str = Field(default="", description="Tool input arguments for tool_use blocks") - tool_output: str = Field(default="", description="Tool output for tool_result blocks") - - # For media blocks - media_url: str = Field(default="", description="URL for image/audio/video blocks") - - @property - def preview(self) -> str: - """Return a short preview of the block content.""" - return self.format(_DEFAULT_MAX_BLOCK_TEXT_PREVIEW_LENGTH) - - def _truncate(self, text: str, max_length: int) -> str: - """Truncate text with ellipsis, replacing newlines with spaces.""" - text = text.replace("\n", " ") - if len(text) <= max_length: - return text - return text[:max_length] + "..." - - # pylint: disable=too-many-return-statements - def format(self, max_length: int = _DEFAULT_MAX_FORMATTER_TEXT_LENGTH, include_thinking: bool = True) -> str: - """Format block content to string representation. - - Args: - max_length: Maximum length of text content in the output. - include_thinking: Whether to include thinking block content. - - Returns: - Formatted string representation of the block. - """ - if self.block_type == "text": - if not self.text: - return "" - return f"[text]: {self._truncate(self.text, max_length)}" - if self.block_type == "thinking": - if not include_thinking or not self.text: - return "" - return f"[think]: {self._truncate(self.text, max_length)}" - if self.block_type in ("image", "audio", "video"): - content = self.media_url if self.media_url else "" - return f"[{self.block_type}]: {content}" - if self.block_type == "tool_use": - content = f"{self.tool_name} params={self._truncate(self.tool_input, max_length)}" - return f"[tool_use]: {content}" - if self.block_type == "tool_result": - if not self.tool_output: - return "" - display_output = self.tool_output.split(_TRUNCATION_NOTICE_MARKER)[0] - content = f"{self.tool_name} output={self._truncate(display_output, max_length)}" - return f"[tool_result]: {content}" - return "" - - -class AsMsgStat(BaseModel): - """Statistics and metadata for a complete AgentScope message.""" - - name: str = Field(default=...) - role: str = Field(default="") - content: list[AsBlockStat] = Field(default_factory=list) - timestamp: str = Field(default="") - metadata: dict = Field(default_factory=dict) - - @property - def total_tokens(self) -> int: - """Return the total token count across all content blocks.""" - return sum(block.token_count for block in self.content) - - @property - def preview(self) -> str: - """Return a short preview of the message content.""" - return self.format(_DEFAULT_MAX_BLOCK_TEXT_PREVIEW_LENGTH) - - def format(self, max_length: int = _DEFAULT_MAX_FORMATTER_TEXT_LENGTH, include_thinking: bool = True) -> str: - """Format message to string representation.""" - time_str = f"[{self.timestamp}] " if self.timestamp else "" - header = f"{time_str}{self.name or self.role}:" - blocks = [block.format(max_length, include_thinking) for block in self.content] - return "\n".join([header] + [b for b in blocks if b]) diff --git a/reme/core/schema/cut_point_result.py b/reme/core/schema/cut_point_result.py deleted file mode 100644 index 34df64fb..00000000 --- a/reme/core/schema/cut_point_result.py +++ /dev/null @@ -1,21 +0,0 @@ -"""Cut point result schemas for context window management.""" - -from pydantic import BaseModel, Field - -from .message import Message - - -class CutPointResult(BaseModel): - """Cut point detection result for conversation compaction.""" - - messages_to_summarize: list[Message] = Field(default_factory=list, description="Complete turns before cut point") - turn_prefix_messages: list[Message] = Field(default_factory=list, description="Turn prefix if split turn") - left_messages: list[Message] = Field(default_factory=list, description="Messages to keep from cut point onwards") - - is_split_turn: bool = Field(default=False, description="Whether cut point is mid-turn") - cut_index: int = Field(default=0, description="Index of cut point in original message list") - - needs_compaction: bool = Field(default=False, description="Whether compaction is actually needed") - token_count: int = Field(default=0, description="Total token count of original messages") - threshold: int = Field(default=0, description="Token threshold that triggers compaction") - accumulated_tokens: int = Field(default=0, description="Tokens accumulated when finding cut point") diff --git a/reme/core/schema/file_metadata.py b/reme/core/schema/file_metadata.py deleted file mode 100644 index 672e9c27..00000000 --- a/reme/core/schema/file_metadata.py +++ /dev/null @@ -1,15 +0,0 @@ -"""File metadata schema.""" - -from pydantic import BaseModel, Field - - -class FileMetadata(BaseModel): - """File metadata with optional extended fields for various use cases.""" - - hash: str = Field(default=..., description="Hash of the file content") - mtime_ms: float = Field(default=..., description="Last modification time in milliseconds") - size: int = Field(default=..., description="File size in bytes") - path: str | None = Field(default=None, description="Relative path to the session file") - content: str | None = Field(default=None, description="Parsed content from the session file") - chunk_count: int | None = Field(default=None, description="Number of chunks in the file") - metadata: dict = Field(default_factory=dict, description="Additional metadata") diff --git a/reme/core/schema/memory_chunk.py b/reme/core/schema/memory_chunk.py deleted file mode 100644 index 36959da1..00000000 --- a/reme/core/schema/memory_chunk.py +++ /dev/null @@ -1,19 +0,0 @@ -"""Memory chunk schema.""" - -from pydantic import BaseModel, Field - -from ..enumeration import MemorySource - - -class MemoryChunk(BaseModel): - """A chunk of memory content with metadata.""" - - id: str = Field(..., description="Unique identifier for the chunk") - path: str = Field(..., description="File path relative to workspace") - source: MemorySource = Field(..., description="Source of the memory data") - start_line: int = Field(..., description="Starting line number in the source file") - end_line: int = Field(..., description="Ending line number in the source file") - text: str = Field(..., description="Text content of the chunk") - hash: str = Field(..., description="Hash of the chunk content") - embedding: list[float] | None = Field(default=None, description="Vector embedding of the chunk") - metadata: dict = Field(default_factory=dict, description="Additional metadata") diff --git a/reme/core/schema/memory_node.py b/reme/core/schema/memory_node.py deleted file mode 100644 index 8889f18f..00000000 --- a/reme/core/schema/memory_node.py +++ /dev/null @@ -1,243 +0,0 @@ -"""Memory schema module for the ReMe AI system. - -This module defines the MemoryNode class for storing and retrieving -memories in the ReMe system. -""" - -import datetime -import hashlib -from typing import Any - -from pydantic import BaseModel, Field, model_validator - -from .vector_node import VectorNode -from ..enumeration import MemoryType - - -def get_now_time() -> str: - """Get current timestamp in YYYY-MM-DD HH:MM:SS format. - - Returns: - str: Current timestamp string in format 'YYYY-MM-DD HH:MM:SS'. - """ - return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") - - -# Length of the memory ID (first N characters of SHA-256 hash) -MEMORY_ID_LENGTH: int = 16 - - -class MemoryNode(BaseModel): - """Memory node for storing memories in the ReMe system. - - Attributes: - memory_id: Unique identifier, auto-generated from content hash. - memory_type: Type of memory (e.g., SUMMARY, PERSONAL). - memory_target: Target or topic this memory relates to. - when_to_use: Condition description for vector retrieval. - content: Actual memory content. - message_time: Time of the message that generated this memory. - ref_memory_id: Reference to related raw history memory. - time_created: Creation timestamp. - time_modified: Last modification timestamp. - author: Author or source of this memory. - score: Relevance or importance score. - vector: Vector embedding of the memory content. - metadata: Additional metadata for extensibility. - """ - - memory_id: str = Field(default="", description="Unique memory identifier") - memory_type: MemoryType = Field(default=..., description="Type of memory") - memory_target: str = Field(default="", description="Target or topic of the memory") - when_to_use: str = Field(default="", description="Condition description for vector retrieval") - content: str = Field(default="", description="Actual memory content") - message_time: str = Field(default="", description="Time of the message that generated this memory") - ref_memory_id: str = Field(default="", description="Reference to related raw history memory ID") - - time_created: str = Field(default_factory=get_now_time, description="Creation timestamp") - time_modified: str = Field(default_factory=get_now_time, description="Last modification timestamp") - author: str = Field(default="", description="Author or source of the memory") - score: float = Field(default=0, description="Relevance or importance score") - - vector: list[float] | None = Field(default=None, description="Vector embedding of the memory content") - metadata: dict[str, Any] = Field(default_factory=dict, description="Additional metadata") - - def _update_modified_time(self) -> "MemoryNode": - """Update time_modified to current timestamp. - - Returns: - Self: Returns self for method chaining. - """ - self.time_modified = get_now_time() - return self - - def _update_memory_id(self) -> "MemoryNode": - """Generate memory_id from SHA-256 hash of content. - - Takes the first MEMORY_ID_LENGTH characters of the hash. - - Returns: - Self: Returns self for method chaining. - """ - if not self.content: - return self - - hash_obj = hashlib.sha256(self.content.encode("utf-8")) - hex_dig = hash_obj.hexdigest() - self.memory_id = hex_dig[:MEMORY_ID_LENGTH] - return self - - @model_validator(mode="after") - def _update_after_init(self) -> "MemoryNode": - """Post-initialization validator. - - Auto-generates memory_id from content if not provided. - - Returns: - Self: Returns self for method chaining. - """ - if not self.memory_id: - self._update_memory_id() - return self - - def __setattr__(self, name: str, value): - """Auto-update timestamps and memory_id when content or when_to_use changes. - - Args: - name: Attribute name being set. - value: New value for the attribute. - """ - should_update: bool = name in ("when_to_use", "content") and getattr(self, name, None) != value - super().__setattr__(name, value) - if should_update: - self._update_modified_time() - if name == "content": - self._update_memory_id() - - def to_vector_node(self) -> VectorNode: - """Convert to VectorNode for vector storage. - - When when_to_use is set, use it as vector content and store content in metadata. - When when_to_use is empty, use content as vector content directly. - - Returns: - VectorNode: Vector node representation of this memory. - """ - safe_metadata: dict[str, str | bool | int | float] = {} - for key, value in self.metadata.items(): - if isinstance(value, (str, bool, int, float)): - safe_metadata[key] = value - elif isinstance(value, (list, tuple, set)): - safe_metadata[key] = ",".join(str(v) for v in value) - else: - safe_metadata[key] = str(value) - - # Build base metadata (shared fields) - metadata: dict[str, Any] = { - "memory_type": self.memory_type.value, - "memory_target": self.memory_target, - "message_time": self.message_time, - "ref_memory_id": self.ref_memory_id, - "time_created": self.time_created, - "time_modified": self.time_modified, - "author": self.author, - "score": self.score, - **safe_metadata, - } - - if self.when_to_use: - # Use when_to_use for vector embedding, store content in metadata - vector_content = self.when_to_use - metadata["content"] = self.content - else: - # Use content directly for vector embedding - vector_content = self.content - - return VectorNode( - vector_id=self.memory_id, - content=vector_content, - vector=self.vector, - metadata=metadata, - ) - - def format( - self, - include_memory_id: bool = True, - include_when_to_use: bool = True, - include_content: bool = True, - include_message_time: bool = True, - ref_memory_id_key: str = "", - ) -> str: - """Format memory node as string with configurable fields.""" - line = "" - - if include_memory_id and self.memory_id: - line += f"memory_id={self.memory_id} " - - if include_message_time and self.message_time: - line += f"[{self.message_time}] " - - if include_when_to_use and self.when_to_use: - line += f"{self.when_to_use} " - - if include_content and self.content: - line += self.content.strip() - - if ref_memory_id_key and self.ref_memory_id: - line += f" {ref_memory_id_key}={self.ref_memory_id}" - - return line.strip() - - @classmethod - def from_vector_node(cls, node: VectorNode) -> "MemoryNode": - """Reconstruct MemoryNode from VectorNode. - - Reverses the to_vector_node conversion: - - If metadata contains 'content': node.content -> when_to_use, metadata['content'] -> content - - Otherwise: node.content -> content, when_to_use remains empty - - Args: - node: VectorNode containing memory data. - - Returns: - Self: Reconstructed MemoryNode instance. - - Raises: - ValueError: If memory_type in metadata is invalid. - """ - metadata = node.metadata.copy() - memory_type_str = metadata.pop("memory_type", None) - - try: - memory_type: MemoryType = MemoryType(memory_type_str) - except ValueError as e: - raise ValueError( - f"Invalid memory_type '{memory_type_str}' in VectorNode metadata. " - f"Valid types are: {[t.value for t in MemoryType]}", - ) from e - - # Restore when_to_use and content based on metadata structure - if "content" in metadata: - # Original had when_to_use set - when_to_use = node.content - content = metadata.pop("content", "") - else: - # Original had empty when_to_use - when_to_use = "" - content = node.content - - return cls( - memory_id=node.vector_id, - memory_type=memory_type, - memory_target=metadata.pop("memory_target", ""), - when_to_use=when_to_use, - content=content, - message_time=metadata.pop("message_time", ""), - ref_memory_id=metadata.pop("ref_memory_id", ""), - time_created=metadata.pop("time_created", ""), - time_modified=metadata.pop("time_modified", ""), - author=metadata.pop("author", ""), - score=metadata.pop("score", 0), - vector=node.vector, - metadata=metadata, - ) diff --git a/reme/core/schema/memory_search_result.py b/reme/core/schema/memory_search_result.py deleted file mode 100644 index dd87dca6..00000000 --- a/reme/core/schema/memory_search_result.py +++ /dev/null @@ -1,23 +0,0 @@ -"""Memory search result schema.""" - -from pydantic import BaseModel, Field - -from ..enumeration import MemorySource - - -class MemorySearchResult(BaseModel): - """Search result from memory index.""" - - path: str = Field(..., description="File path relative to workspace") - start_line: int = Field(..., description="Starting line number of the match") - end_line: int = Field(..., description="Ending line number of the match") - score: float = Field(..., description="Relevance score of the search result") - snippet: str = Field(..., description="Text snippet from the matched content") - source: MemorySource = Field(..., description="Source of the memory data") - raw_metric: float | None = Field(None, description="Raw metric value from search (e.g., distance, rank)") - metadata: dict = Field(default_factory=dict, description="Additional metadata") - - @property - def merge_key(self) -> str: - """Merge key for the search result.""" - return self.path + f":{self.start_line}:{self.end_line}" diff --git a/reme/core/schema/message.py b/reme/core/schema/message.py deleted file mode 100644 index c015eb27..00000000 --- a/reme/core/schema/message.py +++ /dev/null @@ -1,177 +0,0 @@ -"""Data models for multi-modal conversation history and LLM interaction trajectories.""" - -import datetime -import json -import re - -from pydantic import BaseModel, ConfigDict, Field, model_validator - -from .tool_call import ToolCall -from ..enumeration import Role - - -class ContentBlock(BaseModel): - """ - Individual unit of multi-modal content like text, images, or video. - examples: - { - "type": "image_url", - "image_url": { - "url": "https://img.alicdn.com/imgextra/i1/O1CN01gDEY8M1W114Hi3XcN_!!6000000002727-0-tps-1024-406.jpg" - }, - } - - { - "type": "video", - "video": [ - "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/xzsgiz/football1.jpg", - "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/tdescd/football2.jpg", - "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/zefdja/football3.jpg", - "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/aedbqh/football4.jpg", - ], - } - - { - "type": "text", - "text": "How do you solve this problem?" - } - """ - - model_config = ConfigDict(extra="allow") - - type: str = Field(default="") - content: str | dict | list = Field(default="") - - @model_validator(mode="before") - @classmethod - def init_block(cls, data: dict) -> dict: - """Dynamically maps the type-specific key to the content field.""" - content_type = data.get("type", "") - if content_type and content_type in data: - data["content"] = data[content_type] - return data - - def simple_dump(self) -> dict: - """Serializes the block into an API-compatible dictionary format.""" - return { - "type": self.type, - self.type: self.content, - **self.model_extra, - } - - -class Message(BaseModel): - """Data model for a single dialogue entry including roles and tool interactions.""" - - name: str | None = Field(default=None) - role: Role = Field(default=Role.USER) - content: str | list[ContentBlock] = Field(default="") - reasoning_content: str = Field(default="") - tool_calls: list[ToolCall] = Field(default_factory=list) - tool_call_id: str = Field(default="") - time_created: str = Field(default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")) - metadata: dict = Field(default_factory=dict) - - def dump_content(self) -> str | list[dict]: - """Returns content as a raw string or a list of serialized blocks.""" - if isinstance(self.content, str): - return self.content - return [block.simple_dump() for block in self.content] - - def get_text_content(self) -> str: - """Extract plain text content from message, handling both str and list[ContentBlock].""" - if isinstance(self.content, str): - return self.content - return " ".join( - block.content if isinstance(block.content, str) else str(block.content) for block in self.content - ) - - def simple_dump( - self, - add_name: bool = False, - add_reasoning: bool = True, - add_time_created: bool = False, - add_metadata: bool = False, - enable_argument_dict: bool = False, - as_dict: bool = True, - ) -> dict | str: - """Transforms the message into a simplified dictionary for standard APIs.""" - result = {} - if add_name and self.name: - result["name"] = self.name - - result["role"] = self.role.value - result["content"] = self.dump_content() - - if add_reasoning and self.reasoning_content: - result["reasoning_content"] = self.reasoning_content - - if self.tool_calls: - result["tool_calls"] = [ - tc.simple_output_dump( - as_dict=True, - enable_argument_dict=enable_argument_dict, - ) - for tc in self.tool_calls - ] - - if self.tool_call_id: - result["tool_call_id"] = self.tool_call_id - - if add_time_created: - result["time_created"] = self.time_created - - if add_metadata: - result["metadata"] = self.metadata - - return result if as_dict else json.dumps(result, ensure_ascii=False) - - def format_message( - self, - index: int | None = None, - add_time: bool = False, - use_name: bool = False, - add_reasoning: bool = True, - add_tools: bool = True, - strip_markdown_headers: bool = False, - ) -> str: - """Generates a human-readable string representation of the message.""" - prefix = f"round{index} " if index is not None else "" - time_str = f"[{self.time_created}] " if add_time else "" - header = f"{self.name or self.role.value if use_name else self.role.value}:" - - lines = [f"{prefix}{time_str}{header}"] - - def strip_md_func(line): - if strip_markdown_headers: - line = re.sub(r"\n##+ +", "\n", line) - return line - - if add_reasoning and self.reasoning_content: - lines.append(self.reasoning_content) - - if isinstance(self.content, str): - lines.append(strip_md_func(self.content)) - - elif isinstance(self.content, list): - for block in self.content: - text = ( - block.content if isinstance(block.content, str) else json.dumps(block.content, ensure_ascii=False) - ) - text = str(text) - lines.append(strip_md_func(text)) - - if add_tools and self.tool_calls: - for tc in self.tool_calls: - lines.append(f" - tool_call={tc.name} params={tc.arguments}") - - return " ".join(lines).strip() - - -class Trajectory(BaseModel): - """Sequence of messages representing a full conversation session and its evaluation.""" - - task_id: str = Field(default="") - messages: list[Message] = Field(default_factory=list) - score: float = Field(default=0.0) - metadata: dict = Field(default_factory=dict) diff --git a/reme/core/schema/request.py b/reme/core/schema/request.py deleted file mode 100644 index ece942b9..00000000 --- a/reme/core/schema/request.py +++ /dev/null @@ -1,11 +0,0 @@ -"""Defines the data structure for processing incoming user requests and message history.""" - -from pydantic import Field, BaseModel, ConfigDict - - -class Request(BaseModel): - """Represents a structured request payload containing a query, message list, and metadata.""" - - model_config = ConfigDict(extra="allow") - - metadata: dict = Field(default_factory=dict) diff --git a/reme/core/schema/response.py b/reme/core/schema/response.py deleted file mode 100644 index fe753232..00000000 --- a/reme/core/schema/response.py +++ /dev/null @@ -1,13 +0,0 @@ -"""Defines the standardized data structure for model output responses.""" - -from typing import Any - -from pydantic import Field, BaseModel - - -class Response(BaseModel): - """Represents a structured response containing the execution result, status, and metadata.""" - - answer: str | Any = Field(default="") - success: bool = Field(default=True) - metadata: dict = Field(default_factory=dict) diff --git a/reme/core/schema/service_config.py b/reme/core/schema/service_config.py deleted file mode 100644 index 1f8a6816..00000000 --- a/reme/core/schema/service_config.py +++ /dev/null @@ -1,145 +0,0 @@ -"""Configuration schemas for service components using Pydantic models.""" - -import os - -from pydantic import BaseModel, Field, ConfigDict - -from .tool_call import ToolCall - - -class MCPConfig(BaseModel): - """Configuration for Model Context Protocol transport and network settings.""" - - model_config = ConfigDict(extra="allow") - - transport: str = Field(default="stdio") - host: str = Field(default="0.0.0.0") - port: int = Field(default=8001) - - -class HttpConfig(BaseModel): - """Configuration for the HTTP server interface and connection lifecycle.""" - - model_config = ConfigDict(extra="allow") - - host: str = Field(default="0.0.0.0") - port: int = Field(default=8001) - timeout_keep_alive: int = Field(default=3600) - limit_concurrency: int = Field(default=1000) - - -class CmdConfig(BaseModel): - """Configuration for command-line flow execution parameters.""" - - model_config = ConfigDict(extra="allow") - - flow: str = Field(default="") - - -class OpConfig(BaseModel): - """Configuration for op settings and parameters.""" - - model_config = ConfigDict(extra="allow") - - prompt_dict: dict[str, str] = Field(default_factory=dict) - params: dict = Field(default_factory=dict) - - -class FlowConfig(ToolCall): - """Configuration for workflow execution, caching, and error handling.""" - - model_config = ConfigDict(extra="allow") - - flow_content: str = Field(default="") - stream: bool = Field(default=False) - raise_exception: bool = Field(default=True) - enable_cache: bool = Field(default=False) - cache_path: str = Field(default="cache/flow") - cache_expire_hours: float = Field(default=0.1) - - -class BasicConfig(BaseModel): - """Configuration for basic service settings and parameters.""" - - model_config = ConfigDict(extra="allow") - - backend: str = Field(default="") - - -class ModelConfig(BasicConfig): - """Configuration for model-based services with backend and model name.""" - - model_name: str = Field(default="") - - -class LLMConfig(ModelConfig): - """Configuration for Large Language Model backend and model identification.""" - - -class EmbeddingModelConfig(ModelConfig): - """Configuration for embedding model backends and identity.""" - - -class TokenCounterConfig(ModelConfig): - """Configuration for token counting services and model mapping.""" - - -class StoreConfig(BasicConfig): - """Configuration for storage services with embedding model support.""" - - embedding_model: str = Field(default="default") - - -class VectorStoreConfig(StoreConfig): - """Configuration for vector database storage and associated embeddings.""" - - collection_name: str = Field(default="reme") - - -class FileStoreConfig(StoreConfig): - """Configuration for file store database storage and associated embeddings.""" - - store_name: str = Field(default="reme") - - -class FileWatcherConfig(BasicConfig): - """Configuration for file watcher service.""" - - file_store: str = Field(default="") - watch_paths: list[str] = Field(default_factory=list) - - -class ServiceConfig(BasicConfig): - """Root configuration schema aggregating all service-level settings and components.""" - - app_name: str = Field(default=os.getenv("APP_NAME", "ReMe")) - working_dir: str = Field(default=".reme") - enable_logo: bool = Field(default=True) - language: str = Field(default="") - thread_pool_max_workers: int = Field( - default=16, - description="Number of thread pool workers. Set to -1 to disable thread pool.", - ) - ray_max_workers: int = Field(default=-1) - log_to_console: bool = Field(default=True) - log_to_file: bool = Field(default=True) - disabled_flows: list[str] = Field(default_factory=list) - enabled_flows: list[str] = Field(default_factory=list) - - mcp_servers: dict[str, dict] = Field(default_factory=dict) - mcp: MCPConfig = Field(default_factory=MCPConfig) - http: HttpConfig = Field(default_factory=HttpConfig) - cmd: CmdConfig = Field(default_factory=CmdConfig) - ops: dict[str, OpConfig] = Field(default_factory=dict) - flows: dict[str, FlowConfig] = Field(default_factory=dict) - as_llms: dict[str, BasicConfig] = Field(default_factory=dict) - as_llm_formatters: dict[str, BasicConfig] = Field(default_factory=dict) - as_token_counters: dict[str, BasicConfig] = Field(default_factory=dict) - llms: dict[str, LLMConfig] = Field(default_factory=dict) - embedding_models: dict[str, EmbeddingModelConfig] = Field(default_factory=dict) - vector_stores: dict[str, VectorStoreConfig] = Field(default_factory=dict) - file_stores: dict[str, FileStoreConfig] = Field(default_factory=dict) - token_counters: dict[str, TokenCounterConfig] = Field(default_factory=dict) - file_watchers: dict[str, FileWatcherConfig] = Field(default_factory=dict) - - metadata: dict = Field(default_factory=dict) diff --git a/reme/core/schema/stream_chunk.py b/reme/core/schema/stream_chunk.py deleted file mode 100644 index 764981fd..00000000 --- a/reme/core/schema/stream_chunk.py +++ /dev/null @@ -1,14 +0,0 @@ -"""Defines the data structure for individual data packets in a streaming response.""" - -from pydantic import Field, BaseModel - -from ..enumeration import ChunkEnum - - -class StreamChunk(BaseModel): - """Represents a single chunk of streamed data including its type, content, and completion status.""" - - chunk_type: ChunkEnum = Field(default=ChunkEnum.ANSWER) - chunk: str | dict | list = Field(default="") - done: bool = Field(default=False) - metadata: dict = Field(default_factory=dict) diff --git a/reme/core/schema/tool_call.py b/reme/core/schema/tool_call.py deleted file mode 100644 index 19d7e106..00000000 --- a/reme/core/schema/tool_call.py +++ /dev/null @@ -1,227 +0,0 @@ -"""MCP Tool Schema definitions for recursive JSON Schema representation.""" - -import json -from typing import Any, Union, Optional - -from mcp.types import Tool -from pydantic import BaseModel, ConfigDict, Field, model_validator, field_validator - -from ..enumeration.json_schema_enum import JsonSchemaEnum - - -class ToolAttr(BaseModel): - """Recursive model representing JSON Schema attributes for tool parameters.""" - - model_config = ConfigDict(extra="allow") - - type: str = Field(default=str(JsonSchemaEnum.STRING), description="The data type of the attribute") - description: Optional[str] = Field(default=None, description="Description of the attribute") - required: Optional[list[str]] = Field(default=None, description="Required property names for object types") - properties: Optional[dict[str, "ToolAttr"]] = Field(default=None, description="Child properties for objects") - items: Optional[Union[dict[str, Any], "ToolAttr"]] = Field(default=None, description="Schema for array items") - enum: Optional[list[str]] = Field(default=None, description="Allowed values for the attribute") - - @field_validator("type") - @classmethod - def validate_type_is_valid_enum(cls, v: str) -> str: - """Validates that the provided type string exists within JsonSchemaEnum values.""" - valid_types = [str(e) for e in JsonSchemaEnum] - - if v not in valid_types: - raise ValueError(f"Invalid type: '{v}'. Must be one of {valid_types}") - return v - - def simple_input_dump(self) -> dict: - """Serializes the attribute into a standard JSON Schema dictionary.""" - res: dict = {"type": self.type} - if self.description: - res["description"] = self.description - if self.enum: - res["enum"] = self.enum - - if self.type == "object" and self.properties is not None: - res["properties"] = { - k: v.simple_input_dump() if isinstance(v, ToolAttr) else v for k, v in self.properties.items() - } - if self.required is not None: - res["required"] = self.required - - if self.type == "array" and self.items is not None: - res["items"] = self.items.simple_input_dump() if isinstance(self.items, ToolAttr) else self.items - - return res - - -# Enable recursive type resolution -ToolAttr.model_rebuild() - - -class ToolCall(BaseModel): - """ - Model representing a tool definition and its call structure. - Supports parsing from standard JSON Schema formats and converting to MCP Tool objects. - input: - { - "type": "function", - "function": { - "name": "get_current_weather", - "description": "It is very useful when you want to check the weather of a specified city.", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "Cities or counties, such as Beijing, Hangzhou, Yuhang District, etc.", - } - }, - "required": ["location"] - } - } - } - output: - { - "index": 0, - "id": "call_6596dafa2a6a46f7a217da", - "function": { - "arguments": "{\"location\": \"Beijing\"}", - "name": "get_current_weather" - }, - "type": "function", - } - """ - - index: int = 0 - id: str = "" - type: str = "function" - name: str = "" - description: str = "" - - arguments: str = Field(default="", description="JSON string of tool execution arguments") - - parameters: ToolAttr = Field( - default_factory=lambda: ToolAttr(type="object", properties={}, required=[]), - description="Specification for input parameters", - ) - - @model_validator(mode="before") - @classmethod - def init_tool_call(cls, data: dict) -> dict: - """Initializes the model by parsing tool-specific body data.""" - data = data.copy() - t_type = data.get("type", "function") - body = data.get(t_type, {}) - - # Extract basic metadata - data["name"] = body.get("name", data.get("name", "")) - data["arguments"] = body.get("arguments", data.get("arguments", "")) - data["description"] = body.get("description", data.get("description", "")) - - # Handle parameters mapping - if "parameters" in body: - params = body["parameters"] - # If parameters is already a dict, ensure it matches ToolAttr structure - if isinstance(params, dict): - data["parameters"] = ToolAttr(**params) - - # Handle output mapping (if provided in source) - if "output" in body and isinstance(body["output"], dict): - data["output"] = ToolAttr(**body["output"]) - - return data - - def simple_input_dump(self, as_dict: bool = True) -> dict | str: - """Returns a standardized tool definition dictionary or JSON string. - - Args: - as_dict: If True, returns dict; if False, returns JSON string. - """ - result = { - "type": self.type, - self.type: { - "name": self.name, - "description": self.description, - "parameters": self.parameters.simple_input_dump(), - }, - } - return result if as_dict else json.dumps(result, ensure_ascii=False) - - def simple_output_dump(self, as_dict: bool = True, enable_argument_dict: bool = False) -> dict | str: - """Convert ToolCall to output format dictionary or JSON string for API responses.""" - result = { - "index": self.index, - "id": self.id, - self.type: { - "arguments": self.argument_dict if enable_argument_dict else self.arguments, - "name": self.name, - }, - "type": self.type, - } - return result if as_dict else json.dumps(result, ensure_ascii=False) - - @property - def argument_dict(self) -> dict: - """Parse and return arguments as a dictionary.""" - return json.loads(self.arguments) - - def check_argument(self) -> bool: - """Check if arguments can be parsed as valid JSON.""" - try: - _ = self.argument_dict - return True - except Exception: - return False - - def sanitize_and_check_argument(self) -> bool: - """ - Attempt to sanitize and validate arguments JSON. - Common issues from LLM streaming: - - Extra closing brackets: }]}] -> }] - - Missing closing brackets - - Trailing commas - """ - if not self.arguments or not self.arguments.strip(): - return False - - try: - # First try parsing as-is - _ = json.loads(self.arguments) - return True - except json.JSONDecodeError: - pass - - # Try to fix common issues - sanitized = self.arguments.strip() - - # Remove trailing extra brackets/braces - # Pattern: if it ends with multiple closing chars, try removing extras - while len(sanitized) > 1: - try: - json.loads(sanitized) - self.arguments = sanitized # Update with sanitized version - return True - except json.JSONDecodeError: - # Try removing last character - if sanitized[-1] in "]}": - sanitized = sanitized[:-1].rstrip() - else: - break - - return False - - @classmethod - def from_mcp_tool(cls, tool: Tool) -> "ToolCall": - """Creates a ToolCall instance from an MCP Tool object.""" - # MCP Tool inputSchema maps directly to our parameters ToolAttr - return cls( - name=tool.name, - description=tool.description or "", - parameters=ToolAttr(**tool.inputSchema), - ) - - def to_mcp_tool(self) -> Tool: - """Converts the instance back into an MCP Tool object.""" - return Tool( - name=self.name, - description=self.description, - inputSchema=self.parameters.simple_input_dump(), - ) diff --git a/reme/core/schema/truncation_result.py b/reme/core/schema/truncation_result.py deleted file mode 100644 index 18e8cc1a..00000000 --- a/reme/core/schema/truncation_result.py +++ /dev/null @@ -1,35 +0,0 @@ -"""Truncation result schema for command output truncation.""" - -from typing import Literal - -from pydantic import BaseModel, Field - - -class TruncationResult(BaseModel): - """Result of output truncation operation. - - Attributes: - content: The truncated content - truncated: Whether truncation occurred - total_lines: Total number of lines in original output - output_lines: Number of lines in truncated output - total_bytes: Total bytes in original output - output_bytes: Bytes in truncated output - truncated_by: What caused truncation ('lines' or 'bytes') - last_line_partial: Whether last line was partially truncated - """ - - content: str = Field(description="The truncated content") - truncated: bool = Field(description="Whether truncation occurred") - total_lines: int = Field(description="Total number of lines in original output") - output_lines: int = Field(description="Number of lines in truncated output") - total_bytes: int = Field(description="Total bytes in original output") - output_bytes: int = Field(description="Bytes in truncated output") - truncated_by: Literal["lines", "bytes"] | None = Field( - default=None, - description="What caused truncation ('lines' or 'bytes')", - ) - last_line_partial: bool = Field( - default=False, - description="Whether last line was partially truncated", - ) diff --git a/reme/core/schema/vector_node.py b/reme/core/schema/vector_node.py deleted file mode 100644 index 8ee55649..00000000 --- a/reme/core/schema/vector_node.py +++ /dev/null @@ -1,14 +0,0 @@ -"""Defines the data structure for individual vector embedding nodes within a retrieval system.""" - -from uuid import uuid4 - -from pydantic import BaseModel, Field - - -class VectorNode(BaseModel): - """Represents a discrete unit of text content paired with its corresponding vector embedding and metadata.""" - - vector_id: str = Field(default_factory=lambda: uuid4().hex) - content: str = Field(default="") - vector: list[float] | None = Field(default=None) - metadata: dict[str, str | bool | int | float] = Field(default_factory=dict) diff --git a/reme/core/service/__init__.py b/reme/core/service/__init__.py deleted file mode 100644 index ba4e8c40..00000000 --- a/reme/core/service/__init__.py +++ /dev/null @@ -1,18 +0,0 @@ -"""service""" - -from .base_service import BaseService -from .cmd_service import CmdService -from .http_service import HttpService -from .mcp_service import MCPService -from ..registry_factory import R - -__all__ = [ - "BaseService", - "CmdService", - "HttpService", - "MCPService", -] - -R.services.register("cmd")(CmdService) -R.services.register("http")(HttpService) -R.services.register("mcp")(MCPService) diff --git a/reme/core/service/base_service.py b/reme/core/service/base_service.py deleted file mode 100644 index 65d06060..00000000 --- a/reme/core/service/base_service.py +++ /dev/null @@ -1,47 +0,0 @@ -"""Base service definitions for flow management.""" - -from abc import ABC, abstractmethod -from typing import TYPE_CHECKING - -from loguru import logger -from pydantic import BaseModel - -from ..flow import BaseFlow -from ..schema import ToolCall -from ..utils import create_pydantic_model - -if TYPE_CHECKING: - from ..application import Application - - -class BaseService(ABC): - """Abstract base class for services that integrate and execute flows.""" - - def __init__(self, app: "Application", **kwargs): - """Initialize the base service.""" - self.app: "Application" = app - self.service_context = self.app.service_context - self.service_config = self.service_context.service_config - self.kwargs = kwargs - - @abstractmethod - def integrate_flow(self, flow: BaseFlow) -> str | None: - """Integrate a flow into the service and return its name if successful.""" - - @staticmethod - def _prepare_route(flow: BaseFlow) -> tuple[ToolCall, type[BaseModel]]: - """Generate the request model and route name for a flow.""" - tool_call = flow.tool_call - model = create_pydantic_model(tool_call.name, tool_call.parameters) - return tool_call, model - - def run(self): - """Initialize and integrate all flows registered in the global context.""" - flow_names: list[str] = [] - for flow in self.service_context.flows.values(): - flow_name = self.integrate_flow(flow) - if flow_name: - flow_names.append(flow_name) - - if flow_names: - logger.info(f"Integrated {','.join(flow_names)}") diff --git a/reme/core/service/cmd_service.py b/reme/core/service/cmd_service.py deleted file mode 100644 index 4d43b5c3..00000000 --- a/reme/core/service/cmd_service.py +++ /dev/null @@ -1,42 +0,0 @@ -"""Command service module for managing and executing command-based workflows.""" - -from loguru import logger - -from .base_service import BaseService -from ..flow import CmdFlow, BaseFlow -from ..utils import run_coro_safely - - -class CmdService(BaseService): - """Service implementation for handling command flow execution logic.""" - - def __init__(self, **kwargs): - """Initialize the command service instance.""" - super().__init__(**kwargs) - self._cmd_flow: CmdFlow | None = None - - def integrate_flow(self, flow: BaseFlow) -> str | None: - """Integrate the workflow configuration into the command service.""" - self._cmd_flow = CmdFlow(flow=self.service_config.cmd.flow, service_context=self.service_context) - return self._cmd_flow.tool_call.name if self._cmd_flow else None - - def run(self): - """Execute the command flow in either asynchronous or synchronous mode.""" - super().run() - if not self._cmd_flow: - logger.warning("No command flow configured, skipping execution") - return - kwargs = self.service_config.cmd.model_extra - if self._cmd_flow.async_mode: - - async def async_run(): - await self.service_context.start() - return await self._cmd_flow.call(**kwargs) - - response = run_coro_safely(async_run()) - else: - run_coro_safely(self.service_context.start()) - response = self._cmd_flow.call_sync(**kwargs) - - if response.answer: - logger.info(f"response.answer={response.answer}") diff --git a/reme/core/service/http_service.py b/reme/core/service/http_service.py deleted file mode 100644 index c634a11d..00000000 --- a/reme/core/service/http_service.py +++ /dev/null @@ -1,93 +0,0 @@ -"""HTTP service implementation using FastAPI.""" - -import asyncio -from collections.abc import AsyncGenerator -from contextlib import asynccontextmanager - -import uvicorn -from fastapi import FastAPI -from fastapi.middleware.cors import CORSMiddleware -from fastapi.responses import StreamingResponse - -from .base_service import BaseService -from ..flow import BaseFlow -from ..schema import Response -from ..utils import execute_stream_task - - -class HttpService(BaseService): - """Expose flows via HTTP REST and SSE endpoints.""" - - def __init__(self, **kwargs): - """Initialize FastAPI app with CORS and health checks.""" - super().__init__(**kwargs) - - @asynccontextmanager - async def lifespan(_: FastAPI): - await self.app.start() - yield - await self.app.close() - - self.http_service = FastAPI(title=self.service_config.app_name, lifespan=lifespan) - - self.http_service.add_middleware( - CORSMiddleware, - allow_origins=["*"], - allow_credentials=True, - allow_methods=["*"], - allow_headers=["*"], - ) - self.http_service.get("/health")(lambda: {"status": "healthy"}) - - def _integrate_flow(self, flow: BaseFlow) -> str: - """Register a standard flow as a POST endpoint.""" - tool_call, request_model = self._prepare_route(flow) - - async def execute_endpoint(request: request_model) -> Response: - return await flow.call(**request.model_dump(exclude_none=True)) - - self.http_service.post( - path=f"/{tool_call.name}", - response_model=Response, - description=tool_call.description, - )(execute_endpoint) - return tool_call.name - - def _integrate_stream_flow(self, flow: BaseFlow) -> str: - """Register a streaming flow as an SSE endpoint.""" - tool_call, request_model = self._prepare_route(flow) - - async def execute_stream_endpoint(request: request_model) -> StreamingResponse: - stream_queue = asyncio.Queue() - task = asyncio.create_task(flow.call(stream_queue=stream_queue, **request.model_dump(exclude_none=True))) - - async def generate_stream() -> AsyncGenerator[bytes, None]: - async for chunk in execute_stream_task( - stream_queue=stream_queue, - task=task, - task_name=tool_call.name, - output_format="bytes", - ): - yield chunk - - return StreamingResponse(generate_stream(), media_type="text/event-stream") - - self.http_service.post(f"/{tool_call.name}")(execute_stream_endpoint) - return tool_call.name - - def integrate_flow(self, flow: BaseFlow) -> str | None: - """Register a flow based on its streaming configuration.""" - return self._integrate_stream_flow(flow) if flow.stream else self._integrate_flow(flow) - - def run(self): - """Start the Uvicorn server.""" - super().run() - cfg = self.service_config.http - uvicorn.run( - self.http_service, - host=cfg.host, - port=cfg.port, - timeout_keep_alive=cfg.timeout_keep_alive, - limit_concurrency=cfg.limit_concurrency, - **cfg.model_extra, - ) diff --git a/reme/core/service/mcp_service.py b/reme/core/service/mcp_service.py deleted file mode 100644 index 6ed9148c..00000000 --- a/reme/core/service/mcp_service.py +++ /dev/null @@ -1,57 +0,0 @@ -"""Model Context Protocol (MCP) service implementation.""" - -from contextlib import asynccontextmanager - -from fastmcp import FastMCP -from fastmcp.tools import FunctionTool - -from .base_service import BaseService -from ..flow import BaseFlow - - -class MCPService(BaseService): - """Expose flows as Model Context Protocol (MCP) tools.""" - - def __init__(self, **kwargs): - """Initialize FastMCP instance with service settings.""" - super().__init__(**kwargs) - - @asynccontextmanager - async def lifespan(_: FastMCP): - await self.app.start() - yield {} - await self.app.close() - - self.mcp_service = FastMCP(name=self.service_config.app_name, lifespan=lifespan) - - def integrate_flow(self, flow: BaseFlow) -> str | None: - """Register a non-streaming flow as an MCP tool.""" - if flow.stream: - return None - - tool_call, request_model = self._prepare_route(flow) - - async def execute_tool(**kwargs): - """Execute flow logic and return the string answer.""" - request_instance = request_model(**kwargs) - response = await flow.call(**request_instance.model_dump(exclude_none=True)) - return response.answer - - self.mcp_service.add_tool( - FunctionTool( - name=tool_call.name, # noqa - description=tool_call.description, # noqa - fn=execute_tool, - parameters=tool_call.parameters.simple_input_dump(), - ), - ) - return tool_call.name - - def run(self): - """Run the MCP server with specified transport protocol.""" - super().run() - cfg = self.service_config.mcp_service - run_args: dict = {"transport": cfg.transport, "show_banner": False, **cfg.model_extra} - if cfg.transport != "stdio": - run_args.update({"host": cfg.host, "port": cfg.port}) - self.mcp_service.run(**run_args) diff --git a/reme/core/service_context.py b/reme/core/service_context.py deleted file mode 100644 index 5e57c9e4..00000000 --- a/reme/core/service_context.py +++ /dev/null @@ -1,112 +0,0 @@ -"""Service context.""" - -from concurrent.futures import ThreadPoolExecutor -from typing import TYPE_CHECKING - -from loguru import logger - -from .base_dict import BaseDict -from .schema import ServiceConfig -from .utils import PydanticConfigParser - -if TYPE_CHECKING: - from agentscope.model import ChatModelBase - from agentscope.formatter import FormatterBase - from agentscope.token import TokenCounterBase - from .llm import BaseLLM - from .embedding import BaseEmbeddingModel - from .vector_store import BaseVectorStore - from .file_store import BaseFileStore - from .token_counter import BaseTokenCounter - from .flow import BaseFlow - from .file_watcher import BaseFileWatcher - - -class ServiceContext(BaseDict): - """Service context.""" - - def __init__( - self, - *args, - service_config: ServiceConfig | None = None, - parser: type[PydanticConfigParser] | None = None, - working_dir: str | None = None, - config_path: str | None = None, - enable_logo: bool = True, - log_to_console: bool = True, - log_to_file: bool = True, - default_as_llm_config: dict | None = None, - default_as_llm_formatter_config: dict | None = None, - default_as_token_counter_config: dict | None = None, - default_llm_config: dict | None = None, - default_embedding_model_config: dict | None = None, - default_vector_store_config: dict | None = None, - default_file_store_config: dict | None = None, - default_token_counter_config: dict | None = None, - default_file_watcher_config: dict | None = None, - **kwargs, - ): - super().__init__() - - if service_config is None: - parser_class = parser if parser is not None else PydanticConfigParser - parser_instance = parser_class(ServiceConfig) - input_args = [] - if config_path: - input_args.append(f"config={config_path}") - if args: - input_args.extend(args) - - if default_as_llm_config: - self._update_section_config(kwargs, "as_llms", **default_as_llm_config) - if default_as_llm_formatter_config: - self._update_section_config(kwargs, "as_llm_formatters", **default_as_llm_formatter_config) - if default_as_token_counter_config: - self._update_section_config(kwargs, "as_token_counters", **default_as_token_counter_config) - if default_llm_config: - self._update_section_config(kwargs, "llms", **default_llm_config) - if default_embedding_model_config: - self._update_section_config(kwargs, "embedding_models", **default_embedding_model_config) - if default_token_counter_config: - self._update_section_config(kwargs, "token_counters", **default_token_counter_config) - if default_vector_store_config: - self._update_section_config(kwargs, "vector_stores", **default_vector_store_config) - if default_file_store_config: - self._update_section_config(kwargs, "file_stores", **default_file_store_config) - if default_file_watcher_config: - self._update_section_config(kwargs, "file_watchers", **default_file_watcher_config) - - kwargs.update( - { - "enable_logo": enable_logo, - "log_to_console": log_to_console, - "log_to_file": log_to_file, - "working_dir": working_dir, - }, - ) - logger.info(f"update with args: {input_args} kwargs: {kwargs}") - service_config = parser_instance.parse_args(*input_args, **kwargs) - - self.service_config: ServiceConfig = service_config - - self.thread_pool: ThreadPoolExecutor | None = None - self.as_llms: dict[str, "ChatModelBase"] = {} - self.as_llm_formatters: dict[str, "FormatterBase"] = {} - self.as_token_counters: dict[str, "TokenCounterBase"] = {} - self.llms: dict[str, "BaseLLM"] = {} - self.embedding_models: dict[str, "BaseEmbeddingModel"] = {} - self.token_counters: dict[str, "BaseTokenCounter"] = {} - self.vector_stores: dict[str, "BaseVectorStore"] = {} - self.file_stores: dict[str, "BaseFileStore"] = {} - self.file_watchers: dict[str, "BaseFileWatcher"] = {} - self.flows: dict[str, "BaseFlow"] = {} - self.mcp_server_mapping: dict[str, dict] = {} - - @staticmethod - def _update_section_config(config: dict, section_name: str, **kwargs): - """Update a specific section of the service config with new values.""" - if section_name not in config: - config[section_name] = {} - if "default" not in config[section_name]: - config[section_name]["default"] = {} - config[section_name]["default"].update(kwargs) diff --git a/reme/core/token_counter/__init__.py b/reme/core/token_counter/__init__.py deleted file mode 100644 index 6a1bca1d..00000000 --- a/reme/core/token_counter/__init__.py +++ /dev/null @@ -1,16 +0,0 @@ -"""token counter""" - -from .base_token_counter import BaseTokenCounter -from .hf_token_counter import HFTokenCounter -from .openai_token_counter import OpenAITokenCounter -from ..registry_factory import R - -__all__ = [ - "BaseTokenCounter", - "HFTokenCounter", - "OpenAITokenCounter", -] - -R.token_counters.register("base")(BaseTokenCounter) -R.token_counters.register("hf")(HFTokenCounter) -R.token_counters.register("openai")(OpenAITokenCounter) diff --git a/reme/core/token_counter/base_token_counter.py b/reme/core/token_counter/base_token_counter.py deleted file mode 100644 index 4fca5015..00000000 --- a/reme/core/token_counter/base_token_counter.py +++ /dev/null @@ -1,55 +0,0 @@ -"""Token counting utility based on character-type rules.""" - -import math -import re - -from loguru import logger - -from ..schema import Message, ToolCall - - -class BaseTokenCounter: - """A rule-based token counter for Chinese and non-Chinese text.""" - - def __init__(self, model_name: str, **kwargs): - """Initialize with model name and additional parameters.""" - self.model_name = model_name - self.kwargs = kwargs - # Matches Chinese characters including extensions - self._cn_regex = re.compile(r"[\u4e00-\u9fff]") - - def _count_chars(self, text: str) -> tuple[int, int]: - """Count Chinese and other characters in a string.""" - if not text: - return 0, 0 - cn_count = len(self._cn_regex.findall(text)) - return cn_count, len(text) - cn_count - - def count_token( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - **_kwargs, - ) -> int: - """Calculate total tokens using the 1:2 (CN) and 1:4 (Other) rule.""" - cn_total = 0 - ot_total = 0 - logger.info("Calculating tokens using rule-based estimation.") - - # Extract text from messages - segments = [] - for msg in messages: - segments.extend([msg.get_text_content(), msg.reasoning_content]) - - # Extract text from tools - if tools: - for tool in tools: - segments.extend([tool.name, tool.description, tool.arguments]) - - # Process all segments - for text in filter(None, segments): - cn_chars, ot_chars = self._count_chars(text) - cn_total += cn_chars - ot_total += ot_chars - - return math.ceil(cn_total / 2) + math.ceil(ot_total / 4) diff --git a/reme/core/token_counter/hf_token_counter.py b/reme/core/token_counter/hf_token_counter.py deleted file mode 100644 index 4dd072a9..00000000 --- a/reme/core/token_counter/hf_token_counter.py +++ /dev/null @@ -1,80 +0,0 @@ -"""HuggingFace token counting utilities.""" - -import os - -from loguru import logger - -from .base_token_counter import BaseTokenCounter -from ..schema import Message, ToolCall - - -class HFTokenCounter(BaseTokenCounter): - """Token counter using transformers.AutoTokenizer.apply_chat_template.""" - - def __init__( - self, - model_name: str, - use_fast: bool = False, - trust_remote_code: bool = False, - use_mirror: bool = True, - **kwargs, - ): - """Initialize the counter with model config and lazy tokenizer loading.""" - super().__init__(model_name=model_name, **kwargs) - self.use_fast = use_fast - self.trust_remote_code = trust_remote_code - self.use_mirror = use_mirror - self._tokenizer = None - - def _ensure_tokenizer(self): - """Initialize and cache the HuggingFace tokenizer safely.""" - if self._tokenizer: - return self._tokenizer - - if self.use_mirror: - os.environ.setdefault("HF_ENDPOINT", "https://hf-mirror.com") - - try: - from transformers import AutoTokenizer - - logger.info("Initializing HuggingFace tokenizer for {}", self.model_name) - - tokenizer = AutoTokenizer.from_pretrained( - self.model_name, - use_fast=self.use_fast, - trust_remote_code=self.trust_remote_code, - **self.kwargs, - ) - - if not hasattr(tokenizer, "chat_template") or tokenizer.chat_template is None: - raise ValueError(f"Model {self.model_name} lacks a chat template.") - - self._tokenizer = tokenizer - return tokenizer - except Exception as e: - logger.error("Failed to load tokenizer {}: {}", self.model_name, e) - raise - - def count_token( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - **kwargs, - ) -> int: - """Calculate total tokens for messages and tools using the chat template.""" - tokenizer = self._ensure_tokenizer() - - # Serialize inputs for the template - formatted_msgs = [m.simple_dump() for m in messages] - formatted_tools = [t.simple_input_dump() for t in tools] if tools else None - - # Setting tokenize=True and leaving return_tensors=None returns a List[int] - tokens = tokenizer.apply_chat_template( - formatted_msgs, - tools=formatted_tools, - add_generation_prompt=kwargs.pop("add_generation_prompt", False), - tokenize=True, - **kwargs, - ) - - return len(tokens) diff --git a/reme/core/token_counter/openai_token_counter.py b/reme/core/token_counter/openai_token_counter.py deleted file mode 100644 index 0c482145..00000000 --- a/reme/core/token_counter/openai_token_counter.py +++ /dev/null @@ -1,58 +0,0 @@ -"""Token counting implementation for OpenAI-compatible models.""" - -import json - -from loguru import logger - -from .base_token_counter import BaseTokenCounter -from ..schema import Message, ToolCall - - -class OpenAITokenCounter(BaseTokenCounter): - """Token counter for OpenAI models using tiktoken.""" - - def __init__(self, model_name: str, **kwargs): - super().__init__(model_name, **kwargs) - self._encoding = None - - @property - def encoding(self): - """Get or initialize the tiktoken encoding for the specified model.""" - if self._encoding is None: - import tiktoken - - try: - self._encoding = tiktoken.encoding_for_model(self.model_name) - except KeyError: - logger.warning(f"Model {self.model_name} not found; falling back to o200k_base.") - self._encoding = tiktoken.get_encoding("o200k_base") - return self._encoding - - def count_token( - self, - messages: list[Message], - tools: list[ToolCall] | None = None, - **_kwargs, - ) -> int: - """Calculate total tokens for a request including messages and tool definitions.""" - enc = self.encoding - total_tokens = 0 - - for msg in messages: - # Every message has <|start|>{role/name}\n{content}<|end|>\n - total_tokens += 3 # Base overhead per message - if msg.content: - total_tokens += len(enc.encode(msg.get_text_content())) - - if msg.tool_calls: - for tc in msg.tool_calls: - dump = json.dumps(tc.simple_output_dump(), ensure_ascii=False) - total_tokens += len(enc.encode(dump)) - - if tools: - # Account for tool/function definitions if provided - tool_json = json.dumps([t.simple_input_dump() for t in tools], ensure_ascii=False) - total_tokens += len(enc.encode(tool_json)) - - total_tokens += 3 # Every reply is primed with <|start|>assistant<|message|> - return total_tokens diff --git a/reme/core/tools/__init__.py b/reme/core/tools/__init__.py deleted file mode 100644 index 8005793a..00000000 --- a/reme/core/tools/__init__.py +++ /dev/null @@ -1,45 +0,0 @@ -"""tools""" - -from .execute_code import ExecuteCode -from .execute_shell import ExecuteShell - -# file tools -from .file.base_file_tool import BaseFileTool -from .file.bash_tool import BashTool -from .file.edit_tool import EditTool -from .file.find_tool import FindTool -from .file.grep_tool import GrepTool -from .file.ls_tool import LsTool -from .file.read_tool import ReadTool -from .file.write_tool import WriteTool - -# search tools -from .search.dashscope_search import DashscopeSearch -from .search.mock_search import MockSearch -from .search.tavily_search import TavilySearch -from .think_tool import ThinkTool -from ..registry_factory import R - -__all__ = [ - # base tools - "ThinkTool", - "ExecuteCode", - "ExecuteShell", - # file tools - "BaseFileTool", - "BashTool", - "EditTool", - "FindTool", - "GrepTool", - "LsTool", - "ReadTool", - "WriteTool", - # search tools - "DashscopeSearch", - "TavilySearch", - "MockSearch", -] - -for name in __all__: - tool_class = globals()[name] - R.ops.register(tool_class) diff --git a/reme/core/tools/execute_code.py b/reme/core/tools/execute_code.py deleted file mode 100644 index 7631777d..00000000 --- a/reme/core/tools/execute_code.py +++ /dev/null @@ -1,40 +0,0 @@ -"""Code execution tool for running Python code dynamically. - -This module provides an operation that can execute Python code strings -and return the output or error messages. -""" - -from ..op import BaseTool -from ..schema import ToolCall -from ..utils import exec_code, async_exec_code - - -class ExecuteCode(BaseTool): - """Operation for executing Python code dynamically. - - This operation takes Python code as input, executes it in a safe context, - and returns the output or any error messages that occur during execution. - """ - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": self.get_prompt("tool"), - "parameters": { - "type": "object", - "properties": { - "code": { - "type": "string", - "description": "code", - }, - }, - "required": ["code"], - }, - }, - ) - - async def execute(self): - return await async_exec_code(self.context.code) - - def execute_sync(self): - return exec_code(self.context.code) diff --git a/reme/core/tools/execute_shell.py b/reme/core/tools/execute_shell.py deleted file mode 100644 index aa77c036..00000000 --- a/reme/core/tools/execute_shell.py +++ /dev/null @@ -1,46 +0,0 @@ -"""Shell command execution tool. - -This module provides an operation that can execute shell commands -asynchronously and return the output, error, and exit code. -""" - -from ..op import BaseTool -from ..schema import ToolCall -from ..utils import run_shell_command - - -class ExecuteShell(BaseTool): - """Operation for executing shell commands asynchronously. - - This operation takes a shell command as input, executes it asynchronously, - and returns the stdout, stderr, and exit code in a formatted result. - """ - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": self.get_prompt("tool"), - "parameters": { - "type": "object", - "properties": { - "command": { - "type": "string", - "description": "command", - }, - }, - "required": ["command"], - }, - }, - ) - - async def execute(self): - command: str = self.context.command - stdout, stderr, return_code = await run_shell_command(command) - result_parts = [ - f"Command: {command}", - f"Output: {stdout if stdout else '(empty)'}", - f"Error: {stderr if stderr else '(none)'}", - f"Exit Code: {return_code if return_code is not None else '(none)'}", - ] - - return "\n".join(result_parts) diff --git a/reme/core/tools/file/__init__.py b/reme/core/tools/file/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/reme/core/tools/file/base_file_tool.py b/reme/core/tools/file/base_file_tool.py deleted file mode 100644 index e7183f65..00000000 --- a/reme/core/tools/file/base_file_tool.py +++ /dev/null @@ -1,42 +0,0 @@ -"""Base class for file system tools with unified error handling.""" - -from loguru import logger - -from ...op import BaseTool -from ...runtime_context import RuntimeContext - - -class BaseFileTool(BaseTool): - """Base class for file system tools. - - Features: - - No retry logic (max_retries=1) - - Catches all exceptions and returns error messages to LLM - - Simplifies error handling in subclasses - """ - - def __init__(self, **kwargs): - """Initialize fs tool with no retry.""" - kwargs.setdefault("max_retries", 1) - kwargs.setdefault("raise_exception", False) - super().__init__(**kwargs) - - async def call(self, context: RuntimeContext = None, **kwargs): - """Execute the tool with unified error handling. - - This method catches all exceptions and returns error messages - to the LLM instead of raising them. - """ - self.context = RuntimeContext.from_context(context, **kwargs) - - try: - await self.before_execute() - response = await self.execute() - response = await self.after_execute(response) - return response - - except Exception as e: - # Return error message to LLM instead of raising - error_msg = f"{self.__class__.__name__} failed: {str(e)}" - logger.error(error_msg) - return await self.after_execute(error_msg) diff --git a/reme/core/tools/file/bash_tool.py b/reme/core/tools/file/bash_tool.py deleted file mode 100644 index c40c2018..00000000 --- a/reme/core/tools/file/bash_tool.py +++ /dev/null @@ -1,185 +0,0 @@ -"""Bash command execution tool with production-grade features. - -This module provides a production-grade tool for executing bash commands with: -- Smart output truncation (keeps last N lines/bytes to prevent memory issues) -- Process tree termination (prevents orphan processes) -""" - -import asyncio -import os -import platform -import signal -from pathlib import Path - -from .base_file_tool import BaseFileTool -from .truncate import DEFAULT_MAX_BYTES, DEFAULT_MAX_LINES, truncate_tail -from ...schema import ToolCall, TruncationResult - - -def get_shell_config() -> tuple[str, list[str]]: - """Get the appropriate shell and arguments for the current platform. - - Returns: - Tuple of (shell_path, args) for subprocess execution - """ - system = platform.system() - - if system == "Windows": - # Use PowerShell on Windows - return "powershell.exe", ["-Command"] - else: - # Use bash on Unix-like systems - shell = os.environ.get("SHELL", "/bin/bash") - return shell, ["-c"] - - -def kill_process_tree(pid: int) -> None: - """Kill a process and all its children. - - Args: - pid: Process ID to kill - """ - try: - if platform.system() == "Windows": - # Windows: use taskkill - os.system(f"taskkill /F /T /PID {pid}") - else: - # Unix: kill process group - try: - os.killpg(os.getpgid(pid), signal.SIGTERM) - except ProcessLookupError: - pass # Process already dead - except Exception: - pass # Best effort - - -class BashTool(BaseFileTool): - """Production-grade tool for executing bash commands. - - Features: - - Smart output truncation (preserves last N lines or M bytes) - - Kills entire process tree on timeout (prevents orphan processes) - """ - - def __init__(self, cwd: str | None = None, command_prefix: str | None = None): - """Initialize bash tool. - - Args: - cwd: Working directory (defaults to current directory) - command_prefix: Optional prefix prepended to every command - """ - super().__init__() - self.cwd = cwd or os.getcwd() - self.command_prefix = command_prefix - - def _build_tool_call(self) -> ToolCall: - max_kb = DEFAULT_MAX_BYTES // 1024 - return ToolCall( - **{ - "description": ( - f"Execute a bash command in the current working directory. " - f"Returns stdout and stderr. Output is truncated to last " - f"{DEFAULT_MAX_LINES} lines or {max_kb}KB (whichever is hit first). " - f"Optionally provide a timeout in seconds." - ), - "parameters": { - "type": "object", - "properties": { - "command": { - "type": "string", - "description": "Bash command to execute", - }, - "timeout": { - "type": "number", - "description": "Timeout in seconds (optional, no default timeout)", - }, - }, - "required": ["command"], - }, - }, - ) - - async def execute(self) -> str: - """Execute the bash command with production-grade features.""" - command: str = self.context.command - timeout: float | None = self.context.get("timeout", None) - - # Apply command prefix if configured - if self.command_prefix: - command = f"{self.command_prefix}\n{command}" - - # Verify working directory exists - if not Path(self.cwd).exists(): - raise FileNotFoundError( - f"Working directory does not exist: {self.cwd}\n" f"Cannot execute bash commands.", - ) - - # Get shell configuration - shell, shell_args = get_shell_config() - - # Start process - process = await asyncio.create_subprocess_exec( - shell, - *shell_args, - command, - stdout=asyncio.subprocess.PIPE, - stderr=asyncio.subprocess.PIPE, - cwd=self.cwd, - # Create process group for clean termination - preexec_fn=os.setpgrp if platform.system() != "Windows" else None, - ) - - # Execute command with optional timeout - if timeout and timeout > 0: - try: - stdout, stderr = await asyncio.wait_for( - process.communicate(), - timeout=timeout, - ) - except asyncio.TimeoutError as e: - # Kill process tree on timeout - if process.pid: - kill_process_tree(process.pid) - try: - await asyncio.wait_for(process.wait(), timeout=1.0) - except asyncio.TimeoutError: - process.kill() - raise TimeoutError(f"Command timed out after {timeout} seconds") from e - else: - stdout, stderr = await process.communicate() - - # Decode output - full_output = stdout.decode("utf-8", errors="ignore") - if stderr: - stderr_text = stderr.decode("utf-8", errors="ignore") - if full_output: - full_output += "\n" - full_output += stderr_text - - # Apply tail truncation_result to prevent memory issues - truncation_result: TruncationResult = truncate_tail(full_output) - output_text = truncation_result.content or "(no output)" - - # Build truncation_result notice if needed - if truncation_result.truncated: - start_line = truncation_result.total_lines - truncation_result.output_lines + 1 - end_line = truncation_result.total_lines - - if truncation_result.truncated_by == "lines": - output_text += ( - f"\n\n[Output truncated: showing lines {start_line}-{end_line} " - f"of {truncation_result.total_lines} total lines]" - ) - else: - max_kb = DEFAULT_MAX_BYTES // 1024 - output_text += ( - f"\n\n[Output truncated: showing lines {start_line}-{end_line} " - f"of {truncation_result.total_lines} ({max_kb}KB limit reached)]" - ) - - # Handle non-zero exit code - if process.returncode != 0: - output_text += f"\n\nCommand exited with code {process.returncode}" - raise RuntimeError(output_text) - - return output_text diff --git a/reme/core/tools/file/edit_diff.py b/reme/core/tools/file/edit_diff.py deleted file mode 100644 index 76e2fcdb..00000000 --- a/reme/core/tools/file/edit_diff.py +++ /dev/null @@ -1,164 +0,0 @@ -"""Diff utilities for edit tool.""" - -import re -from dataclasses import dataclass -from difflib import unified_diff - - -def detect_line_ending(content: str) -> str: - """Detect line ending style (CRLF or LF).""" - crlf_idx = content.find("\r\n") - lf_idx = content.find("\n") - if lf_idx == -1: - return "\n" - if crlf_idx == -1: - return "\n" - return "\r\n" if crlf_idx < lf_idx else "\n" - - -def normalize_to_lf(text: str) -> str: - """Normalize line endings to LF.""" - return text.replace("\r\n", "\n").replace("\r", "\n") - - -def restore_line_endings(text: str, ending: str) -> str: - """Restore original line endings.""" - return text.replace("\n", ending) if ending == "\r\n" else text - - -def normalize_for_fuzzy_match(text: str) -> str: - """Normalize text for fuzzy matching: strip trailing whitespace, normalize quotes/dashes.""" - lines = text.split("\n") - normalized = "\n".join(line.rstrip() for line in lines) - - # Smart quotes → ASCII - normalized = re.sub(r"[\u2018\u2019\u201A\u201B]", "'", normalized) - normalized = re.sub(r"[\u201C\u201D\u201E\u201F]", '"', normalized) - - # Dashes → hyphen - normalized = re.sub(r"[\u2010\u2011\u2012\u2013\u2014\u2015\u2212]", "-", normalized) - - # Special spaces → regular space - normalized = re.sub(r"[\u00A0\u2002-\u200A\u202F\u205F\u3000]", " ", normalized) - - return normalized - - -@dataclass -class FuzzyMatchResult: - """Result of fuzzy text matching.""" - - found: bool - index: int - match_length: int - used_fuzzy_match: bool - content_for_replacement: str - - -def fuzzy_find_text(content: str, old_text: str) -> FuzzyMatchResult: - """Find old_text in content, trying exact match first, then fuzzy match.""" - # Try exact match - exact_index = content.find(old_text) - if exact_index != -1: - return FuzzyMatchResult( - found=True, - index=exact_index, - match_length=len(old_text), - used_fuzzy_match=False, - content_for_replacement=content, - ) - - # Try fuzzy match - fuzzy_content = normalize_for_fuzzy_match(content) - fuzzy_old_text = normalize_for_fuzzy_match(old_text) - fuzzy_index = fuzzy_content.find(fuzzy_old_text) - - if fuzzy_index == -1: - return FuzzyMatchResult( - found=False, - index=-1, - match_length=0, - used_fuzzy_match=False, - content_for_replacement=content, - ) - - return FuzzyMatchResult( - found=True, - index=fuzzy_index, - match_length=len(fuzzy_old_text), - used_fuzzy_match=True, - content_for_replacement=fuzzy_content, - ) - - -def strip_bom(content: str) -> tuple[str, str]: - """Strip UTF-8 BOM, return (bom, text_without_bom).""" - if content.startswith("\ufeff"): - return "\ufeff", content[1:] - return "", content - - -@dataclass -class DiffResult: - """Result of diff generation.""" - - diff: str - first_changed_line: int | None - - -def generate_diff_string(old_content: str, new_content: str, context_lines: int = 4) -> DiffResult: - """Generate unified diff with line numbers.""" - old_lines = old_content.split("\n") - new_lines = new_content.split("\n") - - # Use difflib to get the changes - diff_lines = list( - unified_diff( - old_lines, - new_lines, - lineterm="", - n=context_lines, - ), - ) - - if not diff_lines: - return DiffResult(diff="", first_changed_line=None) - - # Parse and format the diff - output = [] - first_changed_line = None - max_line_num = max(len(old_lines), len(new_lines)) - line_num_width = len(str(max_line_num)) - - old_line_num = 1 - new_line_num = 1 - - for line in diff_lines[2:]: # Skip header lines - if line.startswith("@@"): - # Parse hunk header - match = re.match(r"@@ -(\d+),?\d* \+(\d+),?\d* @@", line) - if match: - old_line_num = int(match.group(1)) - new_line_num = int(match.group(2)) - continue - - if line.startswith("+"): - if first_changed_line is None: - first_changed_line = new_line_num - line_num = str(new_line_num).rjust(line_num_width) - output.append(f"+{line_num} {line[1:]}") - new_line_num += 1 - elif line.startswith("-"): - if first_changed_line is None: - first_changed_line = new_line_num - line_num = str(old_line_num).rjust(line_num_width) - output.append(f"-{line_num} {line[1:]}") - old_line_num += 1 - else: - # Context line - line_num = str(old_line_num).rjust(line_num_width) - output.append(f" {line_num} {line[1:] if line.startswith(' ') else line}") - old_line_num += 1 - new_line_num += 1 - - return DiffResult(diff="\n".join(output), first_changed_line=first_changed_line) diff --git a/reme/core/tools/file/edit_tool.py b/reme/core/tools/file/edit_tool.py deleted file mode 100644 index 59ba8632..00000000 --- a/reme/core/tools/file/edit_tool.py +++ /dev/null @@ -1,135 +0,0 @@ -"""File editing tool with exact text replacement.""" - -import os -from pathlib import Path - -from .base_file_tool import BaseFileTool -from .edit_diff import ( - detect_line_ending, - fuzzy_find_text, - generate_diff_string, - normalize_for_fuzzy_match, - normalize_to_lf, - restore_line_endings, - strip_bom, -) -from ...schema import ToolCall - - -class EditTool(BaseFileTool): - """Edit a file by replacing exact text.""" - - def __init__(self, cwd: str | None = None): - """Initialize edit tool. - - Args: - cwd: Working directory (defaults to current directory) - """ - super().__init__() - self.cwd = cwd or os.getcwd() - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": ( - "Edit a file by replacing exact text. The oldText must match exactly " - "(including whitespace). Use this for precise, surgical edits." - ), - "parameters": { - "type": "object", - "properties": { - "path": { - "type": "string", - "description": "Path to the file to edit (relative or absolute)", - }, - "oldText": { - "type": "string", - "description": "Exact text to find and replace (must match exactly)", - }, - "newText": { - "type": "string", - "description": "New text to replace the old text with", - }, - }, - "required": ["path", "oldText", "newText"], - }, - }, - ) - - async def execute(self) -> str: - """Execute the edit operation.""" - path: str = self.context.path - old_text: str = self.context.oldText - new_text: str = self.context.newText - - # Resolve path - if not os.path.isabs(path): - absolute_path = os.path.join(self.cwd, path) - else: - absolute_path = path - - # Check file exists and is writable - path_obj = Path(absolute_path) - if not path_obj.exists(): - raise FileNotFoundError(f"File not found: {path}") - - if not os.access(absolute_path, os.R_OK | os.W_OK): - raise PermissionError(f"File not readable/writable: {path}") - - # Read file - with open(absolute_path, "r", encoding="utf-8") as f: - raw_content = f.read() - - # Strip BOM (LLM won't include invisible BOM in oldText) - bom, content = strip_bom(raw_content) - - original_ending = detect_line_ending(content) - normalized_content = normalize_to_lf(content) - normalized_old_text = normalize_to_lf(old_text) - normalized_new_text = normalize_to_lf(new_text) - - # Find old text using fuzzy matching - match_result = fuzzy_find_text(normalized_content, normalized_old_text) - - if not match_result.found: - raise ValueError( - f"Could not find the exact text in {path}. The old text must match " - f"exactly including all whitespace and newlines.", - ) - - # Count occurrences for uniqueness check - fuzzy_content = normalize_for_fuzzy_match(normalized_content) - fuzzy_old_text = normalize_for_fuzzy_match(normalized_old_text) - occurrences = fuzzy_content.count(fuzzy_old_text) - - if occurrences > 1: - raise ValueError( - f"Found {occurrences} occurrences of the text in {path}. " - f"The text must be unique. Please provide more context to make it unique.", - ) - - # Perform replacement - base_content = match_result.content_for_replacement - new_content = ( - base_content[: match_result.index] - + normalized_new_text - + base_content[match_result.index + match_result.match_length :] - ) - - # Verify replacement changed something - if base_content == new_content: - raise ValueError( - f"No changes made to {path}. The replacement produced identical content. " - f"This might indicate an issue with special characters or the text not " - f"exist as expected.", - ) - - # Write file - final_content = bom + restore_line_endings(new_content, original_ending) - with open(absolute_path, "w", encoding="utf-8") as f: - f.write(final_content) - - # Generate diff - diff_result = generate_diff_string(base_content, new_content) - - return f"Successfully replaced text in {path}.\n\n{diff_result.diff}" diff --git a/reme/core/tools/file/find_tool.py b/reme/core/tools/file/find_tool.py deleted file mode 100644 index 73ee285f..00000000 --- a/reme/core/tools/file/find_tool.py +++ /dev/null @@ -1,180 +0,0 @@ -"""File search tool using glob patterns with gitignore support.""" - -import os -from pathlib import Path - -from .base_file_tool import BaseFileTool -from .truncate import FIND_MAX_BYTES, FIND_MAX_LINES, format_size, truncate_head -from ...schema import ToolCall - - -class FindTool(BaseFileTool): - """Search for files by glob pattern, respecting .gitignore.""" - - def __init__(self, cwd: str | None = None): - """Initialize find tool. - - Args: - cwd: Working directory (defaults to current directory) - """ - super().__init__() - self.cwd = cwd or os.getcwd() - - def _build_tool_call(self) -> ToolCall: - max_kb = FIND_MAX_BYTES // 1024 - return ToolCall( - **{ - "description": ( - f"Search for files by glob pattern. Returns matching file paths relative " - f"to the search directory. Respects .gitignore. Output is truncated to " - f"1000 results or {max_kb}KB (whichever is hit first)." - ), - "parameters": { - "type": "object", - "properties": { - "pattern": { - "type": "string", - "description": "Glob pattern to match files, " - "e.g. '*.ts', '**/*.json', or 'src/**/*.spec.ts'", - }, - "path": { - "type": "string", - "description": "Directory to search in (default: current directory)", - }, - "limit": { - "type": "number", - "description": "Maximum number of results (default: 1000)", - }, - }, - "required": ["pattern"], - }, - }, - ) - - def _load_gitignore_patterns(self, search_path: Path) -> list[str]: - """Load gitignore patterns from directory and subdirectories.""" - patterns = ["**/node_modules/**", "**/.git/**"] - - # Load root .gitignore - gitignore_path = search_path / ".gitignore" - if gitignore_path.exists(): - try: - with open(gitignore_path, "r", encoding="utf-8") as f: - for line in f: - line = line.strip() - if line and not line.startswith("#"): - patterns.append(line) - except Exception: - pass # Ignore errors - - # Load nested .gitignore files - try: - for gitignore in search_path.rglob(".gitignore"): - if gitignore == gitignore_path: - continue - try: - with open(gitignore, "r", encoding="utf-8") as f: - for line in f: - line = line.strip() - if line and not line.startswith("#"): - patterns.append(line) - except Exception: - pass # Ignore errors - except Exception: - pass # Ignore glob errors - - return patterns - - def _should_ignore(self, path: Path, ignore_patterns: list[str]) -> bool: - """Check if path matches any ignore pattern.""" - path_str = str(path) - - for pattern in ignore_patterns: - # Simple pattern matching (not full gitignore spec) - if "**" in pattern: - # Recursive match - clean_pattern = pattern.replace("**/", "").replace("/**", "") - if clean_pattern in path_str: - return True - elif "*" in pattern: - # Wildcard match - from fnmatch import fnmatch - - if fnmatch(path.name, pattern): - return True - elif pattern in path_str: - return True - - return False - - async def execute(self) -> str: - """Execute file search.""" - pattern: str = self.context.pattern - search_dir: str = self.context.get("path", ".") - limit: int = self.context.get("limit", 1000) - - # Resolve search path - if not os.path.isabs(search_dir): - search_path = Path(self.cwd) / search_dir - else: - search_path = Path(search_dir) - - # Check if directory exists - if not search_path.exists(): - raise FileNotFoundError(f"Path not found: {search_dir}") - - if not search_path.is_dir(): - raise NotADirectoryError(f"Path is not a directory: {search_dir}") - - # Load gitignore patterns - ignore_patterns = self._load_gitignore_patterns(search_path) - - # Search for files - results = [] - for file_path in search_path.glob(pattern): - if len(results) >= limit: - break - - # Skip if matches ignore patterns - if self._should_ignore(file_path, ignore_patterns): - continue - - # Get relative path - try: - rel_path = file_path.relative_to(search_path) - # Add trailing slash for directories - if file_path.is_dir(): - results.append(f"{rel_path}/") - else: - results.append(str(rel_path)) - except ValueError: - # If relative_to fails, use the path as-is - results.append(str(file_path)) - - # Handle no results - if not results: - return "No files found matching pattern" - - # Sort results for consistency - results.sort() - - # Apply limit and truncation - result_limit_reached = len(results) >= limit - raw_output = "\n".join(results) - truncation = truncate_head(raw_output, max_lines=FIND_MAX_LINES, max_bytes=FIND_MAX_BYTES) - - output = truncation.content - notices = [] - - if result_limit_reached: - notices.append( - f"{limit} results limit reached. Use limit={limit * 2} for more, or refine pattern", - ) - - if truncation.truncated: - notices.append(f"{format_size(FIND_MAX_BYTES)} limit reached") - - if notices: - output += f"\n\n[{'. '.join(notices)}]" - - return output diff --git a/reme/core/tools/file/grep_tool.py b/reme/core/tools/file/grep_tool.py deleted file mode 100644 index a0f1c1ec..00000000 --- a/reme/core/tools/file/grep_tool.py +++ /dev/null @@ -1,273 +0,0 @@ -"""Grep tool for searching file contents using ripgrep. - -This module provides a tool for searching file contents with: -- Pattern matching (regex or literal string) -- Smart output truncation (prevents memory issues) -- Context lines support -- Respects .gitignore -""" - -import asyncio -import json -import os -import shutil -from pathlib import Path - -from .base_file_tool import BaseFileTool -from .truncate import ( - DEFAULT_MAX_BYTES, - GREP_MAX_LINE_LENGTH, - format_size, - truncate_head, - truncate_line, -) -from ...schema import ToolCall - -# Default limits -DEFAULT_LIMIT = 100 # Maximum number of matches - - -class GrepTool(BaseFileTool): - """Tool for searching file contents using ripgrep. - - Features: - - Pattern matching with regex or literal string - - Context lines support - - Smart output truncation - - Respects .gitignore - """ - - def __init__(self, cwd: str | None = None): - """Initialize grep tool. - - Args: - cwd: Working directory (defaults to current directory) - """ - super().__init__() - self.cwd = cwd or os.getcwd() - - def _build_tool_call(self) -> ToolCall: - max_kb = DEFAULT_MAX_BYTES // 1024 - return ToolCall( - **{ - "description": ( - f"Search file contents for a pattern. Returns matching lines with " - f"file paths and line numbers. Respects .gitignore. Output is " - f"truncated to {DEFAULT_LIMIT} matches or {max_kb}KB (whichever is " - f"hit first). Long lines are truncated to {GREP_MAX_LINE_LENGTH} chars." - ), - "parameters": { - "type": "object", - "properties": { - "pattern": { - "type": "string", - "description": "Search pattern (regex or literal string)", - }, - "path": { - "type": "string", - "description": "Directory or file to search (default: current directory)", - }, - "glob": { - "type": "string", - "description": "Filter files by glob pattern, e.g. '*.ts' or '**/*.spec.ts'", - }, - "ignoreCase": { - "type": "boolean", - "description": "Case-insensitive search (default: false)", - }, - "literal": { - "type": "boolean", - "description": "Treat pattern as literal string instead of regex (default: false)", - }, - "contextLines": { - "type": "number", - "description": "Number of lines to show before and after each match (default: 0)", - }, - "limit": { - "type": "number", - "description": f"Maximum number of matches to return (default: {DEFAULT_LIMIT})", - }, - }, - "required": ["pattern"], - }, - }, - ) - - async def execute(self) -> str: - """Execute the grep search.""" - pattern: str = self.context.pattern - search_path: str = self.context.get("path", ".") - glob: str | None = self.context.get("glob", None) - ignore_case: bool = self.context.get("ignoreCase", False) - literal: bool = self.context.get("literal", False) - context_lines: int = self.context.get("contextLines", 0) - limit: int = self.context.get("limit", DEFAULT_LIMIT) - - # Check if ripgrep is available - rg_path = shutil.which("rg") - if not rg_path: - raise RuntimeError( - "ripgrep (rg) is not available. Please install it:\n" - " macOS: brew install ripgrep\n" - " Ubuntu: apt-get install ripgrep\n" - " Other: https://github.com/BurntSushi/ripgrep", - ) - - # Resolve search path - if not os.path.isabs(search_path): - search_path = os.path.join(self.cwd, search_path) - - # Check if path exists - if not Path(search_path).exists(): - raise FileNotFoundError(f"Path not found: {search_path}") - - is_directory = Path(search_path).is_dir() - effective_limit = max(1, limit) - - # Build ripgrep arguments - args = [ - rg_path, - "--json", - "--line-number", - "--color=never", - "--hidden", - ] - - if ignore_case: - args.append("--ignore-case") - - if literal: - args.append("--fixed-strings") - - if glob: - args.extend(["--glob", glob]) - - args.extend([pattern, search_path]) - - # Execute ripgrep - process = await asyncio.create_subprocess_exec( - *args, - stdout=asyncio.subprocess.PIPE, - stderr=asyncio.subprocess.PIPE, - cwd=self.cwd, - ) - - stdout, stderr = await process.communicate() - - # Parse JSON output - matches = [] - match_count = 0 - lines_truncated = False - - for line in stdout.decode("utf-8", errors="ignore").splitlines(): - if not line.strip() or match_count >= effective_limit: - break - - try: - event = json.loads(line) - except json.JSONDecodeError: - continue - - if event.get("type") == "match": - match_count += 1 - data = event.get("data", {}) - file_path = data.get("path", {}).get("text", "") - line_number = data.get("line_number", 0) - - if file_path and line_number: - matches.append({"file_path": file_path, "line_number": line_number}) - - if match_count >= effective_limit: - break - - # Check for errors - if process.returncode not in (0, 1) and match_count == 0: - error_msg = stderr.decode("utf-8", errors="ignore").strip() - if error_msg: - raise RuntimeError(error_msg) - raise RuntimeError(f"ripgrep exited with code {process.returncode}") - - # No matches found - if match_count == 0: - return "No matches found" - - # Format matches with context - output_lines = [] - file_cache = {} - - for match in matches: - file_path = match["file_path"] - line_number = match["line_number"] - - # Read file if not cached - if file_path not in file_cache: - try: - with open(file_path, "r", encoding="utf-8", errors="ignore") as f: - file_cache[file_path] = f.read().replace("\r\n", "\n").replace("\r", "\n").split("\n") - except Exception: - file_cache[file_path] = [] - - lines = file_cache[file_path] - - # Format relative path - if is_directory: - relative_path = os.path.relpath(file_path, search_path) - if not relative_path.startswith(".."): - display_path = relative_path.replace("\\", "/") - else: - display_path = os.path.basename(file_path) - else: - display_path = os.path.basename(file_path) - - # Generate context block - if not lines: - output_lines.append(f"{display_path}:{line_number}: (unable to read file)") - continue - - context_value = max(0, context_lines) - start = max(1, line_number - context_value) if context_value > 0 else line_number - end = min(len(lines), line_number + context_value) if context_value > 0 else line_number - - for current in range(start, end + 1): - if current < 1 or current > len(lines): - continue - - line_text = lines[current - 1] - is_match_line = current == line_number - - # Truncate long lines - truncated_text, was_truncated = truncate_line(line_text) - if was_truncated: - lines_truncated = True - - if is_match_line: - output_lines.append(f"{display_path}:{current}: {truncated_text}") - else: - output_lines.append(f"{display_path}-{current}- {truncated_text}") - - # Apply byte truncation - raw_output = "\n".join(output_lines) - truncation = truncate_head(raw_output, max_lines=999999999) - - output = truncation.content - notices = [] - - # Add notices - if match_count >= effective_limit: - notices.append( - f"{effective_limit} matches limit reached. " - f"Use limit={effective_limit * 2} for more, or refine pattern", - ) - - if truncation.truncated: - notices.append(f"{format_size(DEFAULT_MAX_BYTES)} limit reached") - - if lines_truncated: - notices.append( - f"Some lines truncated to {GREP_MAX_LINE_LENGTH} chars. " f"Use read tool to see full lines", - ) - - if notices: - output += f"\n\n[{'. '.join(notices)}]" - - return output diff --git a/reme/core/tools/file/ls_tool.py b/reme/core/tools/file/ls_tool.py deleted file mode 100644 index e0c6fc3c..00000000 --- a/reme/core/tools/file/ls_tool.py +++ /dev/null @@ -1,125 +0,0 @@ -"""Directory listing tool with truncation support.""" - -import os -from pathlib import Path - -from .base_file_tool import BaseFileTool -from .truncate import DEFAULT_MAX_BYTES, truncate_head -from ...schema import ToolCall - -DEFAULT_LIMIT = 500 - - -class LsTool(BaseFileTool): - """List directory contents with smart truncation. - - Features: - - Returns entries sorted alphabetically (case-insensitive) - - Directory indicators ('/' suffix) - - Includes dotfiles - - Entry count limiting - - Byte truncation - """ - - def __init__(self, cwd: str | None = None): - """Initialize ls tool. - - Args: - cwd: Working directory (defaults to current directory) - """ - super().__init__() - self.cwd = cwd or os.getcwd() - - def _build_tool_call(self) -> ToolCall: - max_kb = DEFAULT_MAX_BYTES // 1024 - return ToolCall( - **{ - "description": ( - f"List directory contents. Returns entries sorted alphabetically, " - f"with '/' suffix for directories. Includes dotfiles. Output is truncated " - f"to {DEFAULT_LIMIT} entries or {max_kb}KB (whichever is hit first)." - ), - "parameters": { - "type": "object", - "properties": { - "path": { - "type": "string", - "description": "Directory to list (default: current directory)", - }, - "limit": { - "type": "number", - "description": f"Maximum number of entries to return (default: {DEFAULT_LIMIT})", - }, - }, - "required": [], - }, - }, - ) - - async def execute(self) -> str: - """List directory contents with production-grade features.""" - path: str | None = self.context.get("path", None) - limit: int | None = self.context.get("limit", None) - - # Resolve directory path - dir_path = Path(self.cwd) / (path or ".") - dir_path = dir_path.resolve() - effective_limit = limit if limit is not None else DEFAULT_LIMIT - - # Check if path exists - if not dir_path.exists(): - raise FileNotFoundError(f"Path not found: {dir_path}") - - # Check if path is a directory - if not dir_path.is_dir(): - raise NotADirectoryError(f"Not a directory: {dir_path}") - - # Read directory entries - entries = list(dir_path.iterdir()) - - # Sort alphabetically (case-insensitive) - entries.sort(key=lambda e: e.name.lower()) - - # Format entries with directory indicators - results: list[str] = [] - entry_limit_reached = False - - for entry in entries: - if len(results) >= effective_limit: - entry_limit_reached = True - break - - try: - # Add '/' suffix for directories - suffix = "/" if entry.is_dir() else "" - results.append(entry.name + suffix) - except Exception: - # Skip entries we can't stat - continue - - # Handle empty directory - if len(results) == 0: - return "(empty directory)" - - # Apply byte truncation - raw_output = "\n".join(results) - truncation_result = truncate_head(raw_output, max_lines=float("inf")) - - output_text = truncation_result.content - - # Build notices - notices: list[str] = [] - - if entry_limit_reached: - notices.append( - f"{effective_limit} entries limit reached. Use limit={effective_limit * 2} for more", - ) - - if truncation_result.truncated: - max_kb = DEFAULT_MAX_BYTES // 1024 - notices.append(f"{max_kb}KB limit reached") - - if notices: - output_text += f"\n\n[{'. '.join(notices)}]" - - return output_text diff --git a/reme/core/tools/file/read_tool.py b/reme/core/tools/file/read_tool.py deleted file mode 100644 index 1c7e8b7e..00000000 --- a/reme/core/tools/file/read_tool.py +++ /dev/null @@ -1,227 +0,0 @@ -"""Read file tool with smart truncation and image support. - -Features: -- Reads text files with offset/limit support -- Detects and handles image files (jpg, png, gif, webp) -- Smart truncation to prevent memory issues -""" - -import os -from pathlib import Path - -from .base_file_tool import BaseFileTool -from .truncate import DEFAULT_MAX_BYTES, DEFAULT_MAX_LINES, format_size, truncate_head -from ...schema import ToolCall - -# Supported image extensions -IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".gif", ".webp"} - - -def is_image_file(path: str) -> bool: - """Check if file is a supported image type. - - Args: - path: File path to check - - Returns: - True if file is a supported image - """ - return Path(path).suffix.lower() in IMAGE_EXTENSIONS - - -class ReadTool(BaseFileTool): - """Read file contents with smart truncation. - - Features: - - Supports text files and images (jpg, png, gif, webp) - - Smart truncation for large files - - Offset/limit for reading specific portions - """ - - def __init__(self, cwd: str | None = None): - """Initialize read tool. - - Args: - cwd: Working directory (defaults to current directory) - """ - super().__init__() - self.cwd = cwd or os.getcwd() - - def _build_tool_call(self) -> ToolCall: - max_kb = DEFAULT_MAX_BYTES // 1024 - return ToolCall( - **{ - "description": ( - f"Read the contents of a file. Supports text files and images " - f"(jpg, png, gif, webp). Images are sent as attachments. For text files, " - f"output is truncated to {DEFAULT_MAX_LINES} lines or {max_kb}KB " - f"(whichever is hit first). Use offset/limit for large files. " - f"When you need the full file, continue with offset until complete." - ), - "parameters": { - "type": "object", - "properties": { - "path": { - "type": "string", - "description": "Path to the file to read (relative or absolute)", - }, - "offset": { - "type": "number", - "description": "Line number to start reading from (1-indexed)", - }, - "limit": { - "type": "number", - "description": "Maximum number of lines to read", - }, - }, - "required": ["path"], - }, - }, - ) - - async def execute(self) -> str: - """Execute the read operation.""" - path: str = self.context.path - offset: int | None = self.context.get("offset", None) - limit: int | None = self.context.get("limit", None) - - # Resolve path - if not os.path.isabs(path): - absolute_path = os.path.join(self.cwd, path) - else: - absolute_path = path - absolute_path = os.path.normpath(absolute_path) - - # Check file exists and is readable - if not os.path.exists(absolute_path): - raise ValueError(f"File not found: {path}") - - if not os.path.isfile(absolute_path): - raise ValueError(f"Not a file: {path}") - - if not os.access(absolute_path, os.R_OK): - raise ValueError(f"File not readable: {path}") - - # Check if image - if is_image_file(absolute_path): - return await self._read_image(absolute_path, path) - else: - return await self._read_text(absolute_path, path, offset, limit) - - @staticmethod - async def _read_image(absolute_path: str, display_path: str) -> str: - """Read and return image file information. - - Args: - absolute_path: Absolute path to image - display_path: Path to display to user - - Returns: - Image information text - """ - # Get file size - file_size = os.path.getsize(absolute_path) - file_ext = Path(absolute_path).suffix.lower() - - # For Python tools, we typically can't return image data directly to LLM - # So we return a descriptive message - return ( - f"Read image file [{file_ext}]\n" - f"Path: {display_path}\n" - f"Size: {format_size(file_size)}\n" - f"Note: Image content cannot be displayed in text format. " - f"Use bash tool or other methods to process the image." - ) - - @staticmethod - async def _read_text( - absolute_path: str, - _display_path: str, - offset: int | None, - limit: int | None, - ) -> str: - """Read text file with smart truncation. - - Args: - absolute_path: Absolute path to file - _display_path: Path to display to user - offset: Starting line (1-indexed) - limit: Maximum lines to read - - Returns: - File contents with truncation notices - """ - # Read file - try: - with open(absolute_path, "r", encoding="utf-8") as f: - content = f.read() - except UnicodeDecodeError: - # Try with error handling for binary files - with open(absolute_path, "r", encoding="utf-8", errors="ignore") as f: - content = f.read() - - all_lines = content.split("\n") - total_file_lines = len(all_lines) - - # Validate and apply offset (1-indexed to 0-indexed) - if offset is not None: - if offset < 1: - raise ValueError(f"offset must be >= 1, got {offset}") - start_line = offset - 1 - else: - start_line = 0 - - start_line_display = start_line + 1 - - # Check offset bounds - if start_line >= total_file_lines: - raise IndexError( - f"Offset {offset} is beyond end of file ({total_file_lines} lines total)", - ) - - # Validate and apply limit - if limit is not None: - if limit <= 0: - raise ValueError(f"limit must be positive, got {limit}") - end_line = min(start_line + limit, total_file_lines) - else: - end_line = total_file_lines - - # Extract selected lines - selected_content = "\n".join(all_lines[start_line:end_line]) - - # Apply truncation - truncation = truncate_head(selected_content) - - # Build output with truncation notices - if truncation.truncated: - # Truncation occurred - end_line_display = start_line_display + truncation.output_lines - 1 - next_offset = end_line_display + 1 - - output_text = truncation.content - - if truncation.truncated_by == "lines": - output_text += ( - f"\n\n[Showing lines {start_line_display}-{end_line_display} " - f"of {total_file_lines}. Use offset={next_offset} to continue.]" - ) - else: - max_kb = DEFAULT_MAX_BYTES // 1024 - output_text += ( - f"\n\n[Showing lines {start_line_display}-{end_line_display} " - f"of {total_file_lines} ({max_kb}KB limit). " - f"Use offset={next_offset} to continue.]" - ) - elif end_line < total_file_lines: - # User limit reached but no truncation - remaining = total_file_lines - end_line - next_offset = end_line + 1 - - output_text = truncation.content - output_text += f"\n\n[{remaining} more lines in file. " f"Use offset={next_offset} to continue.]" - else: - # No truncation or limit - output_text = truncation.content - - return output_text diff --git a/reme/core/tools/file/truncate.py b/reme/core/tools/file/truncate.py deleted file mode 100644 index 22f0c4b8..00000000 --- a/reme/core/tools/file/truncate.py +++ /dev/null @@ -1,209 +0,0 @@ -"""fs utils""" - -from typing import Literal - -from ...schema import TruncationResult - -# Default limits for output truncation -DEFAULT_MAX_LINES = 1000 # Maximum lines to keep for tail truncation -DEFAULT_MAX_BYTES = 30 * 1024 # Maximum bytes to keep (30KB) - -# Find tool limits -FIND_MAX_LINES = 2000 # Maximum lines for find output -FIND_MAX_BYTES = 50 * 1024 # 50KB for find output - -# Grep tool limits -GREP_MAX_LINE_LENGTH = 500 # Maximum line length for grep output - - -def format_size(num_bytes: int) -> str: - """Format byte size in human-readable format. - - Args: - num_bytes: Number of bytes - - Returns: - Formatted string (e.g., "1.5KB", "2.3MB") - """ - if num_bytes < 1024: - return f"{num_bytes}B" - elif num_bytes < 1024 * 1024: - return f"{num_bytes / 1024:.1f}KB" - else: - return f"{num_bytes / (1024 * 1024):.1f}MB" - - -def truncate_line(text: str, max_length: int = GREP_MAX_LINE_LENGTH) -> tuple[str, bool]: - """Truncate a single line if it exceeds max length. - - Args: - text: Line text - max_length: Maximum line length - - Returns: - Tuple of (truncated_text, was_truncated) - """ - if len(text) <= max_length: - return text, False - return text[:max_length] + "...", True - - -def truncate_tail( - text: str, - max_lines: int = DEFAULT_MAX_LINES, - max_bytes: int = DEFAULT_MAX_BYTES, -) -> TruncationResult: - """Truncate text to keep only the tail (last portion). - - Keeps the last N lines or M bytes, whichever is hit first. - This is useful for command outputs where the end is most relevant. - - Args: - text: The text to truncate - max_lines: Maximum number of lines to keep - max_bytes: Maximum bytes to keep - - Returns: - TruncationResult with truncated content and metadata - """ - if not text: - return TruncationResult( - content="", - truncated=False, - total_lines=0, - output_lines=0, - total_bytes=0, - output_bytes=0, - ) - - total_bytes = len(text.encode("utf-8")) - lines = text.split("\n") - total_lines = len(lines) - - # Check if we need to truncate - if total_lines <= max_lines and total_bytes <= max_bytes: - return TruncationResult( - content=text, - truncated=False, - total_lines=total_lines, - output_lines=total_lines, - total_bytes=total_bytes, - output_bytes=total_bytes, - ) - - # Keep last N lines - kept_lines = lines[-max_lines:] if total_lines > max_lines else lines - truncated_by: Literal["lines", "bytes"] = "lines" if total_lines > max_lines else "bytes" - - # Check byte limit on kept lines - kept_text = "\n".join(kept_lines) - kept_bytes = len(kept_text.encode("utf-8")) - - # If still over byte limit, truncate further - last_line_partial = False - if kept_bytes > max_bytes: - truncated_by = "bytes" - # Keep truncating from the start until under byte limit - while kept_lines and len("\n".join(kept_lines).encode("utf-8")) > max_bytes: - kept_lines.pop(0) - - # If still over (single line > max_bytes), truncate the line itself - if kept_lines and len("\n".join(kept_lines).encode("utf-8")) > max_bytes: - last_line = kept_lines[-1] - # Binary search to find how much of last line fits - encoded = last_line.encode("utf-8") - if len(encoded) > max_bytes: - last_line_partial = True - # Take last max_bytes of the line - kept_lines[-1] = encoded[-max_bytes:].decode("utf-8", errors="ignore") - - kept_text = "\n".join(kept_lines) - kept_bytes = len(kept_text.encode("utf-8")) - - return TruncationResult( - content=kept_text, - truncated=True, - total_lines=total_lines, - output_lines=len(kept_lines), - total_bytes=total_bytes, - output_bytes=kept_bytes, - truncated_by=truncated_by, - last_line_partial=last_line_partial, - ) - - -def truncate_head( - text: str, - max_lines: int = FIND_MAX_LINES, - max_bytes: int = FIND_MAX_BYTES, -) -> TruncationResult: - """Truncate text to keep only the head (first portion). - - Keeps the first N lines or M bytes, whichever is hit first. - Suitable for file reads where you want to see the beginning. - - Args: - text: The text to truncate - max_lines: Maximum number of lines to keep - max_bytes: Maximum bytes to keep - - Returns: - TruncationResult with truncated content and metadata - """ - if not text: - return TruncationResult( - content="", - truncated=False, - total_lines=0, - output_lines=0, - total_bytes=0, - output_bytes=0, - ) - - total_bytes = len(text.encode("utf-8")) - lines = text.split("\n") - total_lines = len(lines) - - # Check if no truncation needed - if total_lines <= max_lines and total_bytes <= max_bytes: - return TruncationResult( - content=text, - truncated=False, - total_lines=total_lines, - output_lines=total_lines, - total_bytes=total_bytes, - output_bytes=total_bytes, - ) - - # Collect complete lines that fit - kept_lines = [] - kept_bytes = 0 - truncated_by: Literal["lines", "bytes"] = "lines" - - for i, line in enumerate(lines): - if i >= max_lines: - truncated_by = "lines" - break - - # Calculate bytes for this line (+1 for newline except first line) - line_bytes = len(line.encode("utf-8")) + (1 if i > 0 else 0) - - if kept_bytes + line_bytes > max_bytes: - truncated_by = "bytes" - break - - kept_lines.append(line) - kept_bytes += line_bytes - - kept_text = "\n".join(kept_lines) - final_bytes = len(kept_text.encode("utf-8")) - - return TruncationResult( - content=kept_text, - truncated=True, - total_lines=total_lines, - output_lines=len(kept_lines), - total_bytes=total_bytes, - output_bytes=final_bytes, - truncated_by=truncated_by, - ) diff --git a/reme/core/tools/file/write_tool.py b/reme/core/tools/file/write_tool.py deleted file mode 100644 index 75cf6be5..00000000 --- a/reme/core/tools/file/write_tool.py +++ /dev/null @@ -1,81 +0,0 @@ -"""Write tool for creating and overwriting files. - -This module provides a tool for writing content to files with: -- Automatic parent directory creation -- File overwriting (creates if doesn't exist, overwrites if exists) -- Path resolution (relative to working directory) -""" - -import os - -from .base_file_tool import BaseFileTool -from ...schema import ToolCall - - -class WriteTool(BaseFileTool): - """Tool for writing content to files. - - Features: - - Creates file if it doesn't exist, overwrites if it does - - Automatically creates parent directories - - Supports both relative and absolute paths - """ - - def __init__(self, cwd: str | None = None): - """Initialize write tool. - - Args: - cwd: Working directory (defaults to current directory) - """ - super().__init__() - self.cwd = cwd or os.getcwd() - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": ( - "Write content to a file. Creates the file if it doesn't exist, " - "overwrites if it does. Automatically creates parent directories." - ), - "parameters": { - "type": "object", - "properties": { - "path": { - "type": "string", - "description": "Path to the file to write (relative or absolute)", - }, - "content": { - "type": "string", - "description": "Content to write to the file", - }, - }, - "required": ["path", "content"], - }, - }, - ) - - async def execute(self) -> str: - """Execute the write operation.""" - path: str = self.context.path - content: str = self.context.content - - # Resolve path to absolute - if not os.path.isabs(path): - absolute_path = os.path.join(self.cwd, path) - else: - absolute_path = path - - absolute_path = os.path.normpath(absolute_path) - - # Create parent directories if needed - parent_dir = os.path.dirname(absolute_path) - if parent_dir: - os.makedirs(parent_dir, exist_ok=True) - - # Write the file - with open(absolute_path, "w", encoding="utf-8") as f: - f.write(content) - - # Return success message - content_bytes = len(content.encode("utf-8")) - return f"Successfully wrote {content_bytes} bytes to {path}" diff --git a/reme/core/tools/search/__init__.py b/reme/core/tools/search/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/reme/core/tools/search/dashscope_search.py b/reme/core/tools/search/dashscope_search.py deleted file mode 100644 index da28e251..00000000 --- a/reme/core/tools/search/dashscope_search.py +++ /dev/null @@ -1,108 +0,0 @@ -"""Dashscope web search tool. - -This module provides an operation that uses Alibaba Cloud's Dashscope API -to perform web searches with various search strategies. -""" - -import os -from typing import Literal - -from loguru import logger - -from ...op import BaseTool -from ...schema import ToolCall - - -class DashscopeSearch(BaseTool): - """Operation for performing web searches using Dashscope API. - - This operation uses Alibaba Cloud's Dashscope service to search the web - with support for different search strategies (turbo, max, agent) and - optional role-based prompting. - """ - - def __init__( - self, - model: str = "qwen-plus", # qwen-flash - search_strategy: Literal["turbo", "max", "agent"] = "turbo", # agent only for qwen3-max - enable_role_prompt: bool = True, - **kwargs, - ): - - super().__init__(**kwargs) - self.model: str = model - self.search_strategy: Literal["turbo", "max", "agent"] = search_strategy - self.enable_role_prompt: bool = enable_role_prompt - - # see ref: https://help.aliyun.com/zh/model-studio/web-search - self.api_key = os.getenv("DASHSCOPE_API_KEY", "") - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": self.get_prompt("tool"), - "parameters": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "query", - }, - }, - "required": ["query"], - }, - }, - ) - - async def execute(self): - query: str = self.context.query - if self.enable_cache: - cached_result = self.cache.load(query) - if cached_result: - return cached_result["response_content"] - - if self.enable_role_prompt: - user_query = self.prompt_format("role_prompt", query=query) - else: - user_query = query - logger.info(f"user_query={user_query}") - messages: list = [{"role": "user", "content": user_query}] - - import dashscope - - response = await dashscope.AioGeneration.call( - api_key=self.api_key, - model=self.model, - messages=messages, - enable_search=True, - search_options={ - "forced_search": True, - "enable_source": True, - "enable_citation": False, - "search_strategy": self.search_strategy, - }, - result_format="message", - ) - - search_results = [] - response_content = "" - - if response.output: - if response.output.search_info: - search_results = response.output.search_info.get("search_results", []) - - if response.output.choices and len(response.output.choices) > 0: - response_content = response.output.choices[0].message.content - - final_result = { - "query": query, - "search_results": search_results, - "response_content": response_content, - "model": self.model, - "search_strategy": self.search_strategy, - } - - if self.enable_cache: - self.cache.save(query, final_result, expire_hours=self.cache_expire_hours) - - return final_result["response_content"] diff --git a/reme/core/tools/search/mock_search.py b/reme/core/tools/search/mock_search.py deleted file mode 100644 index f2bc36ca..00000000 --- a/reme/core/tools/search/mock_search.py +++ /dev/null @@ -1,62 +0,0 @@ -"""Mock search tool for testing purposes. - -This module provides a mock search operation that generates simulated -search results using an LLM, useful for testing without making actual API calls. -""" - -import json -import random - -from loguru import logger - -from ...enumeration import Role -from ...op import BaseTool -from ...schema import ToolCall, Message -from ...utils import extract_content - - -class MockSearch(BaseTool): - """Operation for generating mock search results. - - This operation generates simulated search results using an LLM, - useful for testing and development without requiring actual search API access. - """ - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": self.get_prompt("tool"), - "parameters": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "query", - }, - }, - "required": ["query"], - }, - }, - ) - - async def execute(self): - query: str = self.context.query - num_results: int = random.randint(0, 5) - messages = [ - Message( - role=Role.SYSTEM, - content="You are a helpful assistant that generates realistic search results in JSON format.", - ), - Message( - role=Role.USER, - content=self.prompt_format("mock_search_prompt", query=query, num_results=num_results), - ), - ] - - logger.info(f"messages={messages}") - - def callback_fn(message: Message): - return extract_content(message.content, "json") - - search_results: str = await self.llm.chat(messages=messages, callback_fn=callback_fn) - return json.dumps(search_results, ensure_ascii=False, indent=2) diff --git a/reme/core/tools/search/tavily_search.py b/reme/core/tools/search/tavily_search.py deleted file mode 100644 index 65b53734..00000000 --- a/reme/core/tools/search/tavily_search.py +++ /dev/null @@ -1,115 +0,0 @@ -"""Tavily web search tool. - -This module provides an operation that uses the Tavily API to perform -web searches and optionally extract content from search results. -""" - -import json -import os - -from loguru import logger - -from ...op import BaseTool -from ...schema import ToolCall - - -class TavilySearch(BaseTool): - """Operation for performing web searches using Tavily API. - - This operation uses the Tavily search service to find web content - and optionally extract raw content from the results, with configurable - character limits for individual items and total content. - """ - - def __init__( - self, - enable_extract: bool = True, - item_max_char_count: int = 20000, - all_max_char_count: int = 50000, - **kwargs, - ): - super().__init__(**kwargs) - self.enable_extract: bool = enable_extract - self.item_max_char_count: int = item_max_char_count - self.all_max_char_count: int = all_max_char_count - self._client = None - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": self.get_prompt("tool"), - "parameters": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "query", - }, - }, - "required": ["query"], - }, - }, - ) - - @property - def client(self): - """Get or create the Tavily async client instance. - - Returns: - AsyncTavilyClient: The Tavily client instance, lazily initialized. - """ - if self._client is None: - from tavily import AsyncTavilyClient - - self._client = AsyncTavilyClient(api_key=os.environ.get("TAVILY_API_KEY", "")) - return self._client - - async def execute(self): - query: str = self.context.query - logger.info(f"tavily_search query={query}") - - if self.enable_cache: - cached_result = self.cache.load(query) - if cached_result: - return json.dumps(cached_result, ensure_ascii=False, indent=2) - - response = await self.client.search(query=query) - logger.info(f"tavily_search response={response}") - - if not self.enable_extract: - if not response.get("results"): - raise RuntimeError("tavily return empty result") - - final_result = {item["url"]: item for item in response["results"]} - - if self.enable_cache and final_result: - self.cache.save(query, final_result, expire_hours=self.cache_expire_hours) - - return json.dumps(final_result, ensure_ascii=False, indent=2) - - url_info_dict = {item["url"]: item for item in response["results"]} - response_extract = await self.client.extract(urls=[item["url"] for item in response["results"]]) - logger.info(f"tavily.response_extract: {response_extract}") - - final_result = {} - all_char_count = 0 - for item in response_extract["results"]: - url = item["url"] - raw_content: str = item["raw_content"] - if len(raw_content) > self.item_max_char_count: - raw_content = raw_content[: self.item_max_char_count] - if all_char_count + len(raw_content) > self.all_max_char_count: - raw_content = raw_content[: self.all_max_char_count - all_char_count] - - if raw_content: - final_result[url] = url_info_dict[url] - final_result[url]["raw_content"] = raw_content - all_char_count += len(raw_content) - - if not final_result: - raise RuntimeError("tavily return empty result") - - if self.enable_cache and final_result: - self.cache.save(query, final_result, expire_hours=self.cache_expire_hours) - - return json.dumps(final_result, ensure_ascii=False, indent=2) diff --git a/reme/core/tools/think_tool.py b/reme/core/tools/think_tool.py deleted file mode 100644 index 828d4043..00000000 --- a/reme/core/tools/think_tool.py +++ /dev/null @@ -1,42 +0,0 @@ -"""Think tool for agent reflection and planning. - -This module provides a tool that prompts the model for explicit reflection -before taking actions, helping agents reason about their next steps. -""" - -from ..op import BaseTool -from ..schema import ToolCall - - -class ThinkTool(BaseTool): - """Utility that prompts the model for explicit reflection text.""" - - def __init__(self, add_output_reflection: bool = False, **kwargs): - """Initialize the think tool.""" - super().__init__(**kwargs) - self.add_output_reflection: bool = add_output_reflection - - def _build_tool_call(self) -> ToolCall: - """Build the tool call schema for think tool.""" - return ToolCall( - **{ - "description": self.get_prompt("tool"), - "parameters": { - "type": "object", - "properties": { - "reflection": { - "type": "string", - "description": self.get_prompt("reflection"), - }, - }, - "required": ["reflection"], - }, - }, - ) - - async def execute(self): - """Execute the think tool by processing reflection input.""" - if self.add_output_reflection: - return self.context["reflection"] - else: - return self.get_prompt("reflection_output") diff --git a/reme/core/utils/__init__.py b/reme/core/utils/__init__.py deleted file mode 100644 index 5ff544bd..00000000 --- a/reme/core/utils/__init__.py +++ /dev/null @@ -1,57 +0,0 @@ -"""utils""" - -from .agentscope_utils import convert_dashscope_to_agentscope -from .cache_handler import CacheHandler -from .case_converter import snake_to_camel, camel_to_snake -from .chunking_utils import chunk_markdown -from .common_utils import run_coro_safely, execute_stream_task, hash_text, cosine_similarity, batch_cosine_similarity -from .env_utils import load_env -from .execute_utils import exec_code, run_shell_command, async_exec_code -from .horse import play_horse_easter_egg -from .http_client import HttpClient -from .llm_utils import extract_content, format_messages, deduplicate_memories -from .logger_utils import init_logger -from .std_logger import get_logger -from .logo_utils import print_logo -from .mcp_client import MCPClient -from .pydantic_config_parser import PydanticConfigParser -from .pydantic_utils import create_pydantic_model -from .singleton import singleton -from .time import timer, get_now_time -from .hf_token_counter_utils import get_hf_token_counter -from .pyseekdb_conn import admin_kwargs_from_client_kwargs, build_pyseekdb_client_kwargs, parse_host_port - -__all__ = [ - "convert_dashscope_to_agentscope", - "CacheHandler", - "snake_to_camel", - "camel_to_snake", - "chunk_markdown", - "run_coro_safely", - "execute_stream_task", - "hash_text", - "cosine_similarity", - "batch_cosine_similarity", - "load_env", - "exec_code", - "async_exec_code", - "run_shell_command", - "play_horse_easter_egg", - "HttpClient", - "extract_content", - "format_messages", - "deduplicate_memories", - "init_logger", - "get_logger", - "print_logo", - "MCPClient", - "PydanticConfigParser", - "create_pydantic_model", - "singleton", - "timer", - "get_now_time", - "get_hf_token_counter", - "admin_kwargs_from_client_kwargs", - "build_pyseekdb_client_kwargs", - "parse_host_port", -] diff --git a/reme/core/utils/agentscope_utils.py b/reme/core/utils/agentscope_utils.py deleted file mode 100644 index 0330f008..00000000 --- a/reme/core/utils/agentscope_utils.py +++ /dev/null @@ -1,389 +0,0 @@ -# -*- coding: utf-8 -*- -"""Utilities for converting between DashScope format and AgentScope Msg format.""" - -import json -from typing import TYPE_CHECKING, Any - -if TYPE_CHECKING: - from agentscope.message import Msg - - -class DashScopeToAgentScopeConverter: - """Converter for DashScope format to AgentScope Msg format.""" - - def __init__(self, default_name: str = "assistant") -> None: - """Initialize the converter. - - Args: - default_name: Default name for assistant messages when not specified. - """ - self.default_name = default_name - - def convert_message( - self, - dashscope_msg: dict[str, Any], - name: str | None = None, - ) -> "Msg": - """Convert a single DashScope format message to AgentScope Msg. - - Args: - dashscope_msg: DashScope format message dictionary containing - 'role', 'content', and optionally 'tool_calls', 'reasoning_content', - 'tool_call_id', 'name'. - name: Override name for the message. If None, uses the name from - dashscope_msg or default_name. - - Returns: - AgentScope Msg object. - - Examples: - >>> converter = DashScopeToAgentScopeConverter() - >>> # Plain text message - >>> ds_msg = {"role": "assistant", "content": "Hello!"} - >>> msg = converter.convert_message(ds_msg) - >>> # Multimodal message with content blocks - >>> ds_msg = { - ... "role": "user", - ... "content": [ - ... {"text": "What's in this image?"}, - ... {"image": "https://example.com/image.jpg"} - ... ] - ... } - >>> msg = converter.convert_message(ds_msg) - >>> # Tool call message - >>> ds_msg = { - ... "role": "assistant", - ... "content": "", - ... "tool_calls": [{ - ... "id": "call_123", - ... "type": "function", - ... "function": { - ... "name": "get_weather", - ... "arguments": '{"city": "Beijing"}' - ... } - ... }] - ... } - >>> msg = converter.convert_message(ds_msg) - """ - from agentscope.message import Msg - - role = dashscope_msg.get("role", "assistant") - if role not in ["user", "assistant", "system"]: - # Map 'tool' role to 'user' since tool results are inputs to assistant - if role == "tool": - role = "user" - else: - role = "assistant" - - # Determine message name - msg_name = name or dashscope_msg.get("name") or (self.default_name if role == "assistant" else role) - - # Handle tool result messages (role="tool") - if dashscope_msg.get("role") == "tool": - content_blocks = self._convert_tool_result_to_blocks(dashscope_msg) - return Msg( - name=msg_name, - content=content_blocks, - role=role, - ) - - # Extract content - raw_content = dashscope_msg.get("content", "") - tool_calls = dashscope_msg.get("tool_calls", []) - reasoning_content = dashscope_msg.get("reasoning_content", "") - - # Check if we need ContentBlocks or plain string - _ = self._has_multimodal_content(raw_content) - has_tools = len(tool_calls) > 0 - has_reasoning = bool(reasoning_content) - - # If only plain text without tools/reasoning/multimodal, use string content - if isinstance(raw_content, str) and not has_tools and not has_reasoning: - return Msg( - name=msg_name, - content=raw_content or "", - role=role, - ) - - # Otherwise, build ContentBlock list - content_blocks = [] - - # Add reasoning content (thinking block) - if has_reasoning: - from agentscope.message import ThinkingBlock - - content_blocks.append( - ThinkingBlock( - type="thinking", - thinking=reasoning_content, - ), - ) - - # Convert content to blocks - content_blocks.extend(self._convert_content_to_blocks(raw_content)) - - # Convert tool calls to blocks - if has_tools: - content_blocks.extend(self._convert_tool_calls_to_blocks(tool_calls)) - - # If we have no blocks but expected to have content, return empty string - if not content_blocks: - return Msg( - name=msg_name, - content="", - role=role, - ) - - return Msg( - name=msg_name, - content=content_blocks, - role=role, - ) - - def convert_messages( - self, - dashscope_msgs: list[dict[str, Any]], - ) -> list["Msg"]: - """Convert a list of DashScope format messages to AgentScope Msgs. - - Args: - dashscope_msgs: List of DashScope format message dictionaries. - - Returns: - List of AgentScope Msg objects. - """ - return [self.convert_message(msg) for msg in dashscope_msgs] - - def _has_multimodal_content(self, content: Any) -> bool: - """Check if content contains multimodal data. - - Args: - content: Content to check (string or list of content blocks). - - Returns: - True if content contains images, audio, or video. - """ - if not isinstance(content, list): - return False - - for item in content: - if isinstance(item, dict): - item_type = item.get("type", "") - if item_type in ["image", "audio", "video", "image_url"]: - return True - # Check for keys that indicate media - if any(key in item for key in ["image", "audio", "video", "image_url"]): - return True - - return False - - def _convert_content_to_blocks( - self, - content: str | list[dict[str, Any]], - ) -> list[Any]: - """Convert DashScope content to AgentScope content blocks. - - Args: - content: DashScope content (string or list of content items). - - Returns: - List of AgentScope content blocks. - """ - from agentscope.message import ( - AudioBlock, - ImageBlock, - TextBlock, - URLSource, - VideoBlock, - ) - - blocks = [] - - if isinstance(content, str): - if content: - blocks.append( - TextBlock( - type="text", - text=content, - ), - ) - elif isinstance(content, list): - for item in content: - if not isinstance(item, dict): - continue - - # Handle text blocks - if "text" in item: - text = item["text"] - if text: - blocks.append( - TextBlock( - type="text", - text=text, - ), - ) - - # Handle image blocks - elif "image" in item or item.get("type") == "image": - url = item.get("image", "") - blocks.append( - ImageBlock( - type="image", - source=URLSource( - type="url", - url=url, - ), - ), - ) - - # Handle image_url format (OpenAI style) - elif "image_url" in item or item.get("type") == "image_url": - image_url = item.get("image_url", {}) - if isinstance(image_url, dict): - url = image_url.get("url", "") - else: - url = str(image_url) - - blocks.append( - ImageBlock( - type="image", - source=URLSource( - type="url", - url=url, - ), - ), - ) - - # Handle audio blocks - elif "audio" in item or item.get("type") == "audio": - url = item.get("audio", "") - blocks.append( - AudioBlock( - type="audio", - source=URLSource( - type="url", - url=url, - ), - ), - ) - - # Handle video blocks - elif "video" in item or item.get("type") == "video": - video_data = item.get("video", "") - # Video can be a URL string or list of frame URLs - if isinstance(video_data, list): - # Use first frame as URL for now - url = video_data[0] if video_data else "" - else: - url = str(video_data) - - blocks.append( - VideoBlock( - type="video", - source=URLSource( - type="url", - url=url, - ), - ), - ) - - return blocks - - def _convert_tool_calls_to_blocks( - self, - tool_calls: list[dict[str, Any]], - ) -> list[Any]: - """Convert DashScope tool_calls to AgentScope ToolUseBlocks. - - Args: - tool_calls: List of DashScope tool call dictionaries. - - Returns: - List of AgentScope ToolUseBlock objects. - """ - from agentscope.message import ToolUseBlock - - blocks = [] - - for tool_call in tool_calls: - tool_id = tool_call.get("id", "") - function = tool_call.get("function", {}) - name = function.get("name", "") - arguments_str = function.get("arguments", "{}") - - # Parse arguments JSON string to dict - try: - arguments = json.loads(arguments_str) - except (json.JSONDecodeError, TypeError): - arguments = {} - - blocks.append( - ToolUseBlock( - type="tool_use", - id=tool_id, - name=name, - input=arguments, - ), - ) - - return blocks - - def _convert_tool_result_to_blocks( - self, - dashscope_msg: dict[str, Any], - ) -> list[Any]: - """Convert DashScope tool result message to AgentScope ToolResultBlock. - - Args: - dashscope_msg: DashScope tool result message with role="tool". - - Returns: - List containing a single ToolResultBlock. - """ - from agentscope.message import ToolResultBlock - - tool_call_id = dashscope_msg.get("tool_call_id", "") - content = dashscope_msg.get("content", "") - name = dashscope_msg.get("name", "") - - # Tool result content should be plain text - return [ - ToolResultBlock( - type="tool_result", - id=tool_call_id, - name=name, - output=content if content else "", - ), - ] - - -def convert_dashscope_to_agentscope( - dashscope_msg: dict[str, Any] | list[dict[str, Any]], - name: str | None = None, - default_name: str = "assistant", -) -> "Msg | list[Msg]": - """Convenience function to convert DashScope format to AgentScope Msg. - - Args: - dashscope_msg: Single message dict or list of message dicts in DashScope format. - name: Override name for the message(s). - default_name: Default name for assistant messages. - - Returns: - Single Msg object or list of Msg objects. - - Examples: - >>> # Single message - >>> msg = convert_dashscope_to_agentscope({"role": "assistant", "content": "Hi"}) - >>> # Multiple messages - >>> msgs = convert_dashscope_to_agentscope([ - ... {"role": "user", "content": "Hello"}, - ... {"role": "assistant", "content": "Hi there!"} - ... ]) - """ - converter = DashScopeToAgentScopeConverter(default_name=default_name) - - if isinstance(dashscope_msg, list): - return converter.convert_messages(dashscope_msg) - else: - return converter.convert_message(dashscope_msg, name=name) diff --git a/reme/core/utils/cache_handler.py b/reme/core/utils/cache_handler.py deleted file mode 100644 index f3b0072f..00000000 --- a/reme/core/utils/cache_handler.py +++ /dev/null @@ -1,198 +0,0 @@ -"""Local file-based cache utility for DataFrames, lists, dicts, and strings.""" - -import json -from datetime import datetime, timedelta -from pathlib import Path -from typing import Any - -import pandas as pd -from loguru import logger - - -class CacheHandler: - """Handles persistent data caching with expiration and type support.""" - - _EXTENSIONS = { - pd.DataFrame: ".csv", - dict: ".json", - list: ".jsonl", - str: ".txt", - } - - _TYPE_NAMES = { - "DataFrame": pd.DataFrame, - "dict": dict, - "list": list, - "str": str, - } - - def __init__(self, cache_dir: str | Path = "cache"): - """Initialize cache directory and load existing metadata.""" - self.cache_dir = Path(cache_dir) - self.cache_dir.mkdir(parents=True, exist_ok=True) - self.metadata_file = self.cache_dir / "metadata.json" - self.metadata: dict[str, Any] = self._load_metadata() - - def set_cache_dir(self, cache_dir: str | Path) -> None: - """Change the cache directory and reload metadata.""" - self.cache_dir = Path(cache_dir) - self.cache_dir.mkdir(parents=True, exist_ok=True) - self.metadata_file = self.cache_dir / "metadata.json" - self.metadata = self._load_metadata() - logger.info(f"Cache directory moved to: {self.cache_dir}") - - def _load_metadata(self) -> dict[str, Any]: - """Load metadata from the JSON file.""" - if self.metadata_file.exists(): - try: - with open(self.metadata_file, "r", encoding="utf-8") as f: - return json.load(f) - except (json.JSONDecodeError, OSError) as e: - logger.warning(f"Metadata load failed: {e}") - return {} - - def _save_metadata(self) -> None: - """Persist metadata to the disk.""" - try: - with open(self.metadata_file, "w", encoding="utf-8") as f: - json.dump(self.metadata, f, ensure_ascii=False, indent=2) - except OSError as e: - logger.error(f"Metadata save failed: {e}") - - def _get_path(self, key: str, data_type: type | None = None) -> Path: - """Resolve the file path based on data type or metadata.""" - ext = ".dat" - if data_type in self._EXTENSIONS: - ext = self._EXTENSIONS[data_type] - elif key in self.metadata: - stored_type = self.metadata[key].get("data_type") - ext = self._EXTENSIONS.get(self._TYPE_NAMES.get(stored_type, None), ".dat") - return self.cache_dir / f"{key}{ext}" - - @staticmethod - def _execute_save(data: Any, path: Path, dtype: type, **kwargs) -> dict: - """Execute type-specific save operations.""" - if dtype is pd.DataFrame: - data.to_csv(path, index=kwargs.get("index", False), encoding="utf-8") - return {"row_count": len(data), "file_size": path.stat().st_size} - - if dtype is dict: - with open(path, "w", encoding="utf-8") as f: - json.dump(data, f, ensure_ascii=False, indent=2) - return {"item_count": len(data), "file_size": path.stat().st_size} - - if dtype is list: - with open(path, "w", encoding="utf-8") as f: - for item in data: - f.write(json.dumps(item, ensure_ascii=False) + "\n") - return {"item_count": len(data), "file_size": path.stat().st_size} - - if dtype is str: - path.write_text(data, encoding=kwargs.get("encoding", "utf-8")) - return {"char_count": len(data), "file_size": path.stat().st_size} - - raise ValueError(f"Unsupported type: {dtype}") - - @staticmethod - def _execute_load(path: Path, type_name: str, **kwargs) -> Any: - """Execute type-specific load operations.""" - if type_name == "DataFrame": - return pd.read_csv(path, encoding=kwargs.get("encoding", "utf-8")) - if type_name == "dict": - with open(path, "r", encoding="utf-8") as f: - return json.load(f) - if type_name == "list": - result = [] - with open(path, "r", encoding="utf-8") as f: - for line in f: - line = line.strip() - if line: - result.append(json.loads(line)) - return result - if type_name == "str": - return path.read_text(encoding=kwargs.get("encoding", "utf-8")) - raise ValueError(f"Unknown data type in metadata: {type_name}") - - def save(self, key: str, data: Any, expire_hours: float | None = None, **kwargs) -> bool: - """Save data to cache with optional expiration.""" - try: - dtype = type(data) - path = self._get_path(key, dtype) - stats = self._execute_save(data, path, dtype, **kwargs) - - now = datetime.now() - self.metadata[key] = { - "created_at": now.isoformat(), - "expire_at": (now + timedelta(hours=expire_hours)).isoformat() if expire_hours else None, - "data_type": dtype.__name__, - **stats, - } - self._save_metadata() - return True - except Exception as e: - logger.error(f"Save failed for {key}: {e}") - return False - - def load(self, key: str, auto_clean: bool = True, **kwargs) -> Any | None: - """Load data from cache if not expired.""" - if self._is_expired(key): - if auto_clean: - self.delete(key) - return None - - path = self._get_path(key) - if not path.exists() or key not in self.metadata: - return None - - try: - return self._execute_load(path, self.metadata[key]["data_type"], **kwargs) - except Exception as e: - logger.error(f"Load failed for {key}: {e}") - return None - - def _is_expired(self, key: str) -> bool: - """Check if the cached entry has expired.""" - entry = self.metadata.get(key) - if not entry or not entry.get("expire_at"): - return False - return datetime.now() > datetime.fromisoformat(entry["expire_at"]) - - def delete(self, key: str) -> bool: - """Remove a specific cache entry and its file.""" - try: - path = self._get_path(key) - if path.exists(): - path.unlink() - if key in self.metadata: - del self.metadata[key] - self._save_metadata() - return True - except OSError as e: - logger.error(f"Delete failed for {key}: {e}") - return False - - def exists(self, key: str) -> bool: - """Check if a valid cache entry exists.""" - return key in self.metadata and not self._is_expired(key) - - def clear_all(self) -> bool: - """Purge all cache files and reset metadata.""" - try: - for file in self.cache_dir.iterdir(): - if file.is_file(): - file.unlink() - self.metadata = {} - self._save_metadata() - return True - except OSError as e: - logger.error(f"Clear all failed: {e}") - return False - - def get_stats(self) -> dict[str, Any]: - """Return cache usage statistics.""" - total_size = sum(f.stat().st_size for f in self.cache_dir.glob("*") if f.is_file()) - return { - "count": len(self.metadata), - "size_mb": round(total_size / (1024 * 1024), 2), - "dir": str(self.cache_dir), - } diff --git a/reme/core/utils/case_converter.py b/reme/core/utils/case_converter.py deleted file mode 100644 index 28815282..00000000 --- a/reme/core/utils/case_converter.py +++ /dev/null @@ -1,28 +0,0 @@ -"""Case conversion utility for PascalCase, camelCase, and snake_case.""" - -import re - -# Acronyms that should remain uppercase in Pascal/camelCase -_ACRONYMS = {"LLM", "API", "URL", "HTTP", "JSON", "XML", "AI", "MCP"} -_ACRONYM_MAP = {word.lower(): word for word in _ACRONYMS} - - -def camel_to_snake(content: str) -> str: - """Convert PascalCase or camelCase to snake_case.""" - # Normalize acronyms to title case (e.g., LLM -> Llm) to assist regex splitting - for word in _ACRONYMS: - content = content.replace(word, word.capitalize()) - - # Insert underscores between case transitions and convert to lowercase - return re.sub(r"(? str: - """Convert snake_case to PascalCase (preserving defined acronyms).""" - return "".join(_ACRONYM_MAP.get(part.lower(), part.capitalize()) for part in content.split("_") if part) - - -if __name__ == "__main__": - # Quick verification - print(camel_to_snake("OpenAILLMClient")) # open_ai_llm_client - print(snake_to_camel("open_ai_llm_client")) # OpenAILLMClient diff --git a/reme/core/utils/chunking_utils.py b/reme/core/utils/chunking_utils.py deleted file mode 100644 index e1511200..00000000 --- a/reme/core/utils/chunking_utils.py +++ /dev/null @@ -1,124 +0,0 @@ -"""Chunking logic for Markdown files.""" - -from .common_utils import hash_text -from ..enumeration import MemorySource -from ..schema import MemoryChunk - - -def chunk_markdown( - text: str, - path: str, - source: MemorySource, - chunk_tokens: int, - overlap: int, -) -> list[MemoryChunk]: - """ - Markdown chunking logic implemented based on the TypeScript version. - - Args: - text: Input text - path: File path - source: Memory source - chunk_tokens: Maximum tokens per chunk - overlap: Overlap tokens between chunks - - Returns: - List of MemoryChunk objects - """ - lines = text.split("\n") - if not lines: - return [] - - # Convert tokens to characters (~1 token = 4 chars) - max_chars = max(32, chunk_tokens * 4) - overlap_chars = max(0, overlap * 4) - - chunks: list[MemoryChunk] = [] - - # Currently building chunk - current: list[dict] = [] # [{'line': str, 'line_no': int}] - current_chars = 0 - - def flush(): - """Add current chunk to results list""" - if not current: - return - - first_entry = current[0] - last_entry = current[-1] - - if not first_entry or not last_entry: - return - - chunk_text = "\n".join([entry["line"] for entry in current]) - start_line = first_entry["line_no"] - end_line = last_entry["line_no"] - - chunk_hash = hash_text(chunk_text) - - chunks.append( - MemoryChunk( - id=hash_text(f"{source}:{path}:{start_line}:{end_line}:{chunk_hash}:{len(chunks)}"), - path=path, - source=source, - start_line=start_line, - end_line=end_line, - text=chunk_text, - hash=chunk_hash, - ), - ) - - def carry_overlap(): - """Keep overlapping part and clear the rest""" - nonlocal current, current_chars - - if overlap_chars <= 0 or not current: - current = [] - current_chars = 0 - return - - acc = 0 - kept = [] - - # Collect lines from the end until reaching overlap size - for j in range(len(current) - 1, -1, -1): - entry = current[j] - if not entry: - continue - - acc += len(entry["line"]) + 1 # +1 for newline - kept.insert(0, entry) # Insert at the beginning to maintain order - - if acc >= overlap_chars: - break - - current = kept - current_chars = sum(len(entry["line"]) + 1 for entry in kept) - - for i, line in enumerate(lines): - line_no = i + 1 - - # Split long lines into multiple segments - segments = [] - if not line: # Empty line - segments.append("") - else: - # If line is too long, split by maximum character count - for start in range(0, len(line), max_chars): - segments.append(line[start : start + max_chars]) - - for segment in segments: - line_size = len(segment) + 1 # +1 for newline - - # If adding current segment would exceed the limit, flush current chunk - if current_chars + line_size > max_chars and current: - flush() - carry_overlap() - - current.append({"line": segment, "line_no": line_no}) - current_chars += line_size - - # Process the final chunk - flush() - - return [c for c in chunks if c.text.strip()] diff --git a/reme/core/utils/common_utils.py b/reme/core/utils/common_utils.py deleted file mode 100644 index 95a3518f..00000000 --- a/reme/core/utils/common_utils.py +++ /dev/null @@ -1,185 +0,0 @@ -"""Common utility functions""" - -import asyncio -import hashlib -from collections.abc import AsyncGenerator, Coroutine -from typing import Any, Literal - -import numpy as np -from loguru import logger - -from ..enumeration import ChunkEnum -from ..schema import StreamChunk - - -def run_coro_safely(coro: Coroutine[Any, Any, Any]) -> Any | asyncio.Task[Any]: - """Run a coroutine in the current event loop or a new one if none exists.""" - try: - # Attempt to retrieve the event loop associated with the current thread - loop = asyncio.get_running_loop() - - except RuntimeError: - # Start a new event loop to run the coroutine to completion - return asyncio.run(coro) - - else: - # Schedule the coroutine as a background task in the active loop - return loop.create_task(coro) - - -async def execute_stream_task( - stream_queue: asyncio.Queue, - task: asyncio.Task, - task_name: str | None = None, - output_format: Literal["str", "bytes", "chunk"] = "str", -) -> AsyncGenerator[str | bytes | StreamChunk, None]: - """ - Core stream flow execution logic. - - Handles streaming from a queue while monitoring the task completion. - Properly manages errors and resource cleanup. - - Args: - stream_queue: Queue to receive StreamChunk objects from - task: Background task executing the flow - task_name: Optional flow name for logging purposes - output_format: Output format control - - "str": SSE-formatted string (default) - - "bytes": SSE-formatted bytes for HTTP responses - - "chunk": Raw StreamChunk objects - - Yields: - - str: SSE-formatted data when output_format="str" - - bytes: SSE-formatted data when output_format="bytes" - - StreamChunk: Raw chunk objects when output_format="chunk" - - Raises: - Exception: Re-raises any exception from the background task - """ - try: - while True: - # Wait for next chunk or check if task failed - get_chunk = asyncio.create_task(stream_queue.get()) - done, _pending = await asyncio.wait({get_chunk, task}, return_when=asyncio.FIRST_COMPLETED) - - # Priority 1: Check if main task finished (may have exception) - if task in done: - # Task finished - check for exceptions first - exc = task.exception() - if exc: - log_msg = f"Task error in {task_name}: {exc}" if task_name else f"Task error: {exc}" - logger.exception(log_msg) - raise exc - - # Task completed successfully - drain remaining chunks if any - if get_chunk in done: - chunk: StreamChunk = get_chunk.result() - if output_format == "chunk": - yield chunk - if chunk.done: - break - else: - if chunk.done: - yield b"data:[DONE]\n\n" if output_format == "bytes" else "data:[DONE]\n\n" - break - data = f"data:{chunk.model_dump_json()}\n\n" - yield data.encode() if output_format == "bytes" else data - else: - # No more chunks, task completed - get_chunk.cancel() - if output_format == "chunk": - yield StreamChunk(chunk_type=ChunkEnum.DONE, chunk="", done=True) - else: - yield b"data:[DONE]\n\n" if output_format == "bytes" else "data:[DONE]\n\n" - break - - elif get_chunk in done: - # Got a chunk from the queue (task still running) - chunk: StreamChunk = get_chunk.result() - - # Handle raw chunk mode - if output_format == "chunk": - yield chunk - if chunk.done: - break - continue - - # Handle SSE format mode (str or bytes) - if chunk.done: - yield b"data:[DONE]\n\n" if output_format == "bytes" else "data:[DONE]\n\n" - break - - data = f"data:{chunk.model_dump_json()}\n\n" - yield data.encode() if output_format == "bytes" else data - - finally: - # Ensure task is cancelled if still running to avoid resource leaks - if not task.done(): - task.cancel() - try: - await task - except asyncio.CancelledError: - pass - - -def hash_text(text: str) -> str: - """Generate SHA-256 hash of text content. - - Args: - text: Input text to hash - - Returns: - Hexadecimal representation of the SHA-256 hash - """ - return hashlib.sha256(text.encode("utf-8")).hexdigest() - - -def cosine_similarity(vec1: list[float], vec2: list[float]) -> float: - """Calculate the cosine similarity between two numeric vectors.""" - if len(vec1) != len(vec2): - raise ValueError(f"Vectors must have same length: {len(vec1)} != {len(vec2)}") - - dot_product = sum(a * b for a, b in zip(vec1, vec2)) - magnitude1 = sum(a * a for a in vec1) ** 0.5 - magnitude2 = sum(b * b for b in vec2) ** 0.5 - - if magnitude1 == 0 or magnitude2 == 0: - return 0.0 - - return dot_product / (magnitude1 * magnitude2) - - -def batch_cosine_similarity(nd_array1: np.ndarray, nd_array2: np.ndarray) -> np.ndarray: - """Calculate cosine similarity matrix between two batches of vectors. - - Args: - nd_array1: Matrix of shape (batch_size1, emb_size) - nd_array2: Matrix of shape (batch_size2, emb_size) - - Returns: - Similarity matrix of shape (batch_size1, batch_size2) where - result[i, j] is the cosine similarity between nd_array1[i] and nd_array2[j] - - Raises: - ValueError: If embedding dimensions don't match - """ - if nd_array1.shape[1] != nd_array2.shape[1]: - raise ValueError(f"Embedding dimensions must match: {nd_array1.shape[1]} != {nd_array2.shape[1]}") - - # Compute dot products: (batch_size1, emb_size) @ (emb_size, batch_size2) - # Result shape: (batch_size1, batch_size2) - dot_products = np.dot(nd_array1, nd_array2.T) - - # Compute L2 norms for each vector - norms1 = np.linalg.norm(nd_array1, axis=1) # Shape: (batch_size1,) - norms2 = np.linalg.norm(nd_array2, axis=1) # Shape: (batch_size2,) - - # Compute outer product of norms: (batch_size1, 1) @ (1, batch_size2) - # Result shape: (batch_size1, batch_size2) - norm_products = np.outer(norms1, norms2) - - # Avoid division by zero - norm_products = np.where(norm_products == 0, 1e-10, norm_products) - - # Compute cosine similarities - return dot_products / norm_products diff --git a/reme/core/utils/env_utils.py b/reme/core/utils/env_utils.py deleted file mode 100644 index b36c7bd5..00000000 --- a/reme/core/utils/env_utils.py +++ /dev/null @@ -1,61 +0,0 @@ -"""Environment variable loader utility for managing .env files.""" - -import os -from pathlib import Path - -from loguru import logger - -# Global flag to ensure environment is loaded only once -_ENV_LOADED = False - - -def _parse_env_file(path: Path) -> None: - """Parse and inject key-value pairs from a .env file into os.environ.""" - try: - with path.open(encoding="utf-8") as file: - for line in file: - line = line.strip() - if not line or line.startswith("#"): - continue - - if "=" in line: - key, value = line.split("=", 1) - # Strip whitespace and common quotes - os.environ[key.strip()] = value.strip().strip("'\"") - except PermissionError as err: - logger.warning(f"Permission denied for {path}: {err}") - except Exception as err: - logger.error(f"Failed to load {path}: {err}") - raise - - -def load_env(path: str | Path | None = None, enable_log: bool = True) -> None: - """Search and load the .env file into the system environment.""" - global _ENV_LOADED # pylint: disable=global-statement - if _ENV_LOADED: - return - - if path: - path = Path(path) - if path.exists(): - _parse_env_file(path) - _ENV_LOADED = True - else: - logger.warning(f".env not found at: {path}") - return - - # Search current directory and up to 5 levels of parents - for directory in [Path.cwd(), *Path.cwd().parents[:5]]: - env_path = directory / ".env" - if env_path.exists(): - if enable_log: - logger.info(f"Loading environment from: {env_path}") - _parse_env_file(env_path) - _ENV_LOADED = True - return - - -def reset_env_flag() -> None: - """Reset the internal load state flag.""" - global _ENV_LOADED # pylint: disable=global-statement - _ENV_LOADED = False diff --git a/reme/core/utils/execute_utils.py b/reme/core/utils/execute_utils.py deleted file mode 100644 index d5333e2f..00000000 --- a/reme/core/utils/execute_utils.py +++ /dev/null @@ -1,142 +0,0 @@ -"""Utility functions for executing code and shell commands. - -This module provides helper functions for running Python code and shell commands, -with support for async execution and output capture. -""" - -import asyncio -import concurrent.futures -import contextlib -from io import StringIO - - -async def run_shell_command(cmd: str, timeout: float | None = 30) -> tuple[str, str, int]: - """Execute a shell command asynchronously. - - Args: - cmd: The shell command to execute. - timeout: Maximum time to wait for command completion in seconds. None for no timeout. - - Returns: - A tuple containing (stdout, stderr, return_code) as strings and integer. - - Raises: - TimeoutError: If the command does not complete within the timeout. - """ - process = await asyncio.create_subprocess_shell( - cmd, - stdout=asyncio.subprocess.PIPE, - stderr=asyncio.subprocess.PIPE, - ) - - try: - if timeout: - stdout, stderr = await asyncio.wait_for(process.communicate(), timeout=timeout) - else: - stdout, stderr = await process.communicate() - except (asyncio.TimeoutError, TimeoutError) as e: - # Kill the child process to avoid orphaned / zombie processes - process.kill() - await process.wait() - raise TimeoutError(f"Shell command timed out after {timeout}s") from e - - return ( - stdout.decode("utf-8", errors="ignore"), - stderr.decode("utf-8", errors="ignore"), - process.returncode, - ) - - -def exec_code( - code: str, - timeout: float | None = 30, - executor: concurrent.futures.ThreadPoolExecutor | None = None, -) -> str: - """Execute Python code and capture the output. - - Args: - code: The Python code string to execute. - timeout: Maximum time to wait for execution in seconds. None for no timeout. - executor: Optional thread pool executor to use. If None, a temporary - single-thread executor is created (and shut down after the call). - Pass a shared executor to amortize thread-creation overhead across - multiple calls. - - Returns: - The captured stdout output, or the error message if execution fails. - - Raises: - TimeoutError: If execution exceeds the timeout. - """ - - def _run() -> str: - redirected_output = StringIO() - with contextlib.redirect_stdout(redirected_output): - exec(code) - return redirected_output.getvalue() - - def _submit(pool: concurrent.futures.ThreadPoolExecutor) -> str: - future = pool.submit(_run) - return future.result(timeout=timeout) - - try: - if executor is not None: - return _submit(executor) - with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool: - return _submit(pool) - - except concurrent.futures.TimeoutError as e: - raise TimeoutError(f"Code execution timed out after {timeout}s") from e - - except Exception as e: - return str(e) - - except BaseException as e: - return str(e) - - -async def async_exec_code( - code: str, - timeout: float | None = 30, - executor: concurrent.futures.ThreadPoolExecutor | None = None, -) -> str: - """Execute Python code asynchronously and capture the output. - - Runs the code in a thread executor to avoid blocking the event loop, - with async-friendly timeout via ``asyncio.wait_for``. - - Args: - code: The Python code string to execute. - timeout: Maximum time to wait for execution in seconds. None for no timeout. - executor: Optional thread pool executor. If None, the default event-loop - executor is used. - - Returns: - The captured stdout output, or the error message if execution fails. - - Raises: - TimeoutError: If execution exceeds the timeout. - """ - - def _run() -> str: - redirected_output = StringIO() - with contextlib.redirect_stdout(redirected_output): - exec(code) - return redirected_output.getvalue() - - loop = asyncio.get_running_loop() - - try: - coro = loop.run_in_executor(executor, _run) - if timeout is not None: - return await asyncio.wait_for(coro, timeout=timeout) - return await coro - - except asyncio.TimeoutError as e: - raise TimeoutError(f"Code execution timed out after {timeout}s") from e - - except Exception as e: - return str(e) - - except BaseException as e: - return str(e) diff --git a/reme/core/utils/hf_token_counter_utils.py b/reme/core/utils/hf_token_counter_utils.py deleted file mode 100644 index a8ab348c..00000000 --- a/reme/core/utils/hf_token_counter_utils.py +++ /dev/null @@ -1,23 +0,0 @@ -"""Utility functions for working with text.""" - -from agentscope.token import HuggingFaceTokenCounter - -_token_counter = None - - -def get_hf_token_counter( - pretrained_model_name_or_path="Qwen/Qwen2.5-7B-Instruct", - use_mirror=True, - use_fast=True, - trust_remote_code=True, -): - """Get or initialize the global token counter instance.""" - global _token_counter - if _token_counter is None: - _token_counter = HuggingFaceTokenCounter( - pretrained_model_name_or_path=pretrained_model_name_or_path, - use_mirror=use_mirror, - use_fast=use_fast, - trust_remote_code=trust_remote_code, - ) - return _token_counter diff --git a/reme/core/utils/horse.py b/reme/core/utils/horse.py deleted file mode 100644 index 3231c8c7..00000000 --- a/reme/core/utils/horse.py +++ /dev/null @@ -1,165 +0,0 @@ -"""Horse Easter egg: fireworks, galloping horse animation, and a blessing.""" - -import math -import random -import shutil -import sys -import time - - -def _mirror_frame(frame: str) -> str: - """Mirror ASCII art horizontally.""" - mirror_map = str.maketrans(r"()/\<>[]{}", r")(\/><][}{") - lines = frame.split("\n") - max_len = max(len(line) for line in lines) if lines else 0 - mirrored = [] - for line in lines: - padded = line.ljust(max_len) - reversed_line = padded[::-1].translate(mirror_map) - mirrored.append(reversed_line) - return "\n".join(mirrored) - - -def play_horse_easter_egg() -> None: - """Play the /horse Easter egg: fireworks, galloping horse, and a blessing.""" - cols = shutil.get_terminal_size((80, 24)).columns - rows = shutil.get_terminal_size((80, 24)).lines - - # -- Fireworks animation (~4 seconds at 8 fps = 32 frames) -- - firework_colors = [ - "\033[91m", # red - "\033[93m", # yellow - "\033[92m", # green - "\033[96m", # cyan - "\033[95m", # magenta - "\033[94m", # blue - ] - reset = "\033[0m" - particles_chars = ["*", ".", "o", "+", "x", "'", "`"] - - class Firework: - """A single firework burst with radial particles.""" - - def __init__(self, cx: int, cy: int, color: str, birth: int): - self.cx = cx - self.cy = cy - self.color = color - self.birth = birth - self.num = random.randint(12, 20) - self.angles = [random.uniform(0, 2 * math.pi) for _ in range(self.num)] - self.speeds = [random.uniform(0.5, 1.5) for _ in range(self.num)] - self.chars = [random.choice(particles_chars) for _ in range(self.num)] - - def particles(self, frame: int): - """Return list of (x, y, char) particle positions for the given frame.""" - age = frame - self.birth - if age < 0 or age > 10: - return [] - pts = [] - for i in range(self.num): - r = self.speeds[i] * age - px = self.cx + int(r * math.cos(self.angles[i]) * 2) # *2 for aspect ratio - py = self.cy + int(r * math.sin(self.angles[i])) - if 0 <= px < cols and 0 <= py < rows - 1: - pts.append((px, py, self.chars[i])) - return pts - - # Hide cursor - sys.stdout.write("\033[?25l") - sys.stdout.flush() - - try: - fireworks: list[Firework] = [] - total_frames = 32 - for f in range(total_frames): - # Spawn new fireworks periodically - if f % 4 == 0: - cx = random.randint(10, cols - 10) - cy = random.randint(2, rows // 2) - color = random.choice(firework_colors) - fireworks.append(Firework(cx, cy, color, f)) - - # Build frame buffer (blank) - buf: dict[tuple[int, int], tuple[str, str]] = {} - for fw in fireworks: - for px, py, ch in fw.particles(f): - buf[(px, py)] = (fw.color, ch) - - # Render - sys.stdout.write("\033[H\033[2J") # clear screen - for y in range(rows - 1): - line_parts: list[str] = [] - x = 0 - for x_pos in sorted(px for (px, py) in buf if py == y): - if x_pos >= x: - line_parts.append(" " * (x_pos - x)) - color, ch = buf[(x_pos, y)] - line_parts.append(f"{color}{ch}{reset}") - x = x_pos + 1 - sys.stdout.write("".join(line_parts) + "\n") - sys.stdout.flush() - time.sleep(1 / 8) # 8 fps - - # Prune old fireworks - fireworks.clear() - - # -- Horse ASCII art (bold yellow) -- - frame_1 = r""" - >>\. - /_ )`. - / _)`^)`. _.---. - (_,' \ `^---- `. - | | - \ / - / \ /___ / \ - / / | \ \ | -""" - frame_2 = r""" - >>\. - /_ )`. - / _)`^)`. _.---. - (_,' \ `^---- `. - | | - \ / - // / ___ / | - / / / | \ \ | -""" - - mirrored_1 = _mirror_frame(frame_1) - mirrored_2 = _mirror_frame(frame_2) - - bold_yellow = "\033[1;33m" - sys.stdout.write("\033[H\033[2J") # clear - - # Short galloping animation (8 cycles) - for i in range(8): - sys.stdout.write("\033[H\033[2J") - horse = mirrored_1 if i % 2 == 0 else mirrored_2 - indent = " " * (i * 3) - print("\n" * 3) - for line in horse.split("\n"): - if line.strip(): - print(f"{bold_yellow}{indent}{line}{reset}") - print(f"{bold_yellow}{'-' * min(i * 3 + 40, cols - 1)}{reset}") - sys.stdout.flush() - time.sleep(0.2) - - # -- Random blessing -- - blessings = [ - ("\u9a6c\u5230\u6210\u529f", "Succeed immediately"), - ("\u9f99\u9a6c\u7cbe\u795e", "Full of vitality"), - ("\u4e07\u9a6c\u5954\u817e", "Thousands of horses galloping"), - ("\u9a6c\u4e0d\u505c\u8e44", "Never stop striving"), - ("\u5feb\u9a6c\u52a0\u97ad", "Full speed ahead"), - ("\u4e00\u9a6c\u5f53\u5148", "Take the lead"), - ] - cn, en = random.choice(blessings) - print() - print(f"{bold_yellow} {cn} - {en}{reset}") - print(f"{bold_yellow} Happy Year of the Horse 2026!{reset}") - print() - - finally: - # Restore cursor - sys.stdout.write("\033[?25h") - sys.stdout.flush() diff --git a/reme/core/utils/http_client.py b/reme/core/utils/http_client.py deleted file mode 100644 index 99c94306..00000000 --- a/reme/core/utils/http_client.py +++ /dev/null @@ -1,90 +0,0 @@ -"""Asynchronous HTTP client for executing flows with built-in retry logic.""" - -import json -from collections.abc import AsyncIterator - -import httpx -from loguru import logger - -from ..schema import Response - - -class HttpClient: - """Async client for flow endpoints with automated retries and error handling.""" - - def __init__( - self, - base_url: str = "http://localhost:8001", - timeout: float = 3600.0, - max_retries: int = 3, - raise_exception: bool = True, - ): - """Initialize the client with base configuration.""" - self.base_url = base_url.rstrip("/") - self.timeout = timeout - self.max_retries = max_retries - self.raise_exception = raise_exception - self.client = httpx.AsyncClient(timeout=timeout) - - async def __aenter__(self): - """Enter async context manager.""" - return self - - async def __aexit__(self, exc_type, exc_val, exc_tb): - """Exit async context manager and close connection.""" - await self.close() - - async def close(self): - """Close the underlying HTTP client.""" - await self.client.aclose() - - async def health_check(self) -> dict[str, str]: - """Check the health status of the flow service.""" - response = await self.client.get(f"{self.base_url}/health") - response.raise_for_status() - return response.json() - - async def execute_flow(self, flow_name: str, **kwargs) -> Response | None: - """Execute a flow with automated retry logic.""" - endpoint = f"{self.base_url}/{flow_name}" - - for attempt in range(self.max_retries): - try: - response = await self.client.post(endpoint, json=kwargs) - response.raise_for_status() - return Response(**response.json()) - - except (httpx.HTTPError, Exception) as e: - logger.error(f"Flow {flow_name} failed (attempt {attempt + 1}/{self.max_retries}): {e}") - if attempt == self.max_retries - 1 and self.raise_exception: - raise e - return None - - async def list_endpoints(self) -> dict: - """Retrieve available endpoints from OpenAPI specification.""" - response = await self.client.get(f"{self.base_url}/openapi.json") - response.raise_for_status() - return response.json() - - async def execute_stream_flow(self, flow_name: str, **kwargs) -> AsyncIterator[dict[str, str]]: - """Execute a flow and yield parsed SSE stream chunks.""" - endpoint = f"{self.base_url}/{flow_name}" - - async with self.client.stream("POST", endpoint, json=kwargs) as response: - response.raise_for_status() - async for line in response.aiter_lines(): - if not line or not line.startswith("data:"): - continue - - content = line.removeprefix("data:").strip() - if content == "[DONE]": - break - - try: - data = json.loads(content) - yield { - "type": data.get("chunk_type", "answer"), - "content": data.get("chunk", ""), - } - except json.JSONDecodeError: - continue diff --git a/reme/core/utils/llm_utils.py b/reme/core/utils/llm_utils.py deleted file mode 100644 index e828c24e..00000000 --- a/reme/core/utils/llm_utils.py +++ /dev/null @@ -1,255 +0,0 @@ -"""Utility functions for processing and formatting LLM-related message data.""" - -import json -import re - -from agentscope.message import Msg -from loguru import logger - -from ..enumeration import Role -from ..schema import Message, Trajectory, MemoryNode, ToolCall - - -def convert_as_msg_to_message(msg) -> Message: - """Convert an agentscope Msg object to the project's Message type.""" - role_str = getattr(msg, "role", "user") - role = ( - Role(role_str.lower()) - if isinstance(role_str, str) and role_str.lower() in [r.value for r in Role] - else Role.USER - ) - - content_blocks = msg.get_content_blocks() - content = "" - reasoning_content = "" - tool_calls = [] - tool_call_id = "" - - for block in content_blocks: - block_type = block["type"] - if block_type == "thinking": - reasoning_content = block["thinking"] - elif block_type == "tool_use": - try: - tool_calls.append( - ToolCall( - id=block["id"], - name=block["name"], - arguments=json.dumps(block["input"], ensure_ascii=False), - ), - ) - except (json.JSONDecodeError, TypeError): - pass - elif block_type == "tool_result": - role = Role.TOOL - tool_call_id = block["id"] - content = block["output"][0]["text"] - else: - content = block[block_type] - - return Message( - name=getattr(msg, "name", None), - role=role, - content=content, - reasoning_content=reasoning_content, - tool_calls=tool_calls, - tool_call_id=tool_call_id, - time_created=getattr(msg, "timestamp", "") or "", - metadata=getattr(msg, "metadata", {}) or {}, - ) - - -def format_messages( - messages: list[Message | dict], - add_index: bool = True, - add_time: bool = True, - use_name: bool = True, - add_reasoning: bool = True, - add_tools: bool = True, - strip_markdown_headers: bool = True, - enable_system: bool = False, -) -> str: - """Formats a list of messages into a single string, optionally filtering system roles.""" - formatted_lines = [] - for i, message in enumerate(messages): - if isinstance(message, dict): - message = Message(**message) - if isinstance(message, Msg): - message = convert_as_msg_to_message(message) - if not enable_system and message.role is Role.SYSTEM: - continue - - formatted_lines.append( - message.format_message( - index=i if add_index else None, - add_time=add_time, - use_name=use_name, - add_reasoning=add_reasoning, - add_tools=add_tools, - strip_markdown_headers=strip_markdown_headers, - ), - ) - return "\n".join(formatted_lines) - - -def merge_messages_content(messages: list[Message | dict]) -> str: - """Merge messages content into a formatted string representation. - - This function processes a list of messages (either Message objects or dicts) - and formats them into a structured string. Different message roles are - formatted differently: - - ASSISTANT: Includes reasoning content, main content, and tool calls - - USER: Includes the user content - - TOOL: Includes tool call results - - Each message is prefixed with a step number (starting from 0) to indicate - its position in the conversation sequence. - - Args: - messages: List of Message objects or dictionaries to merge. If a dict - is provided, it will be converted to a Message object. - - Returns: - Formatted string representation of all messages with step numbers. - Each message is separated by newlines and includes role information. - - Example: - ```python - messages = [ - Message(role=Role.USER, content="What's the weather?"), - Message(role=Role.ASSISTANT, content="Let me check", - tool_calls=[ToolCall(name="get_weather", arguments={})]) - ] - result = merge_messages_content(messages) - # Returns formatted string with step numbers and role information - ``` - """ - content_collector = [] - for i, message in enumerate(messages): - if isinstance(message, dict): - message = Message(**message) - - if message.role is Role.ASSISTANT: - line = ( - f"### step.{i} role={message.role.value} content=\n{message.reasoning_content}\n\n{message.content}\n" - ) - if message.tool_calls: - for tool_call in message.tool_calls: - line += f" - tool call={tool_call.name}\n params={tool_call.arguments}\n" - content_collector.append(line) - - elif message.role is Role.USER: - line = f"### step.{i} role={message.role.value} content=\n{message.content}\n" - content_collector.append(line) - - elif message.role is Role.TOOL: - line = f"### step.{i} role={message.role.value} tool call result=\n{message.content}\n" - content_collector.append(line) - - return "\n".join(content_collector) - - -def parse_json_experience_response(response: str) -> list[dict]: - """Parse JSON formatted experience response""" - try: - # Extract JSON blocks - json_pattern = r"```json\s*([\s\S]*?)\s*```" - json_blocks = re.findall(json_pattern, response) - - if json_blocks: - parsed = json.loads(json_blocks[0]) - - # Handle array format - if isinstance(parsed, list): - experiences = [] - for exp_data in parsed: - if isinstance(exp_data, dict) and ( - ("when_to_use" in exp_data and "experience" in exp_data) - or ("condition" in exp_data and "experience" in exp_data) - ): - experiences.append(exp_data) - - return experiences - - # Handle single object - elif isinstance(parsed, dict) and ( - ("when_to_use" in parsed and "experience" in parsed) - or ("condition" in parsed and "experience" in parsed) - ): - return [parsed] - - # Fallback: try to parse entire response - parsed = json.loads(response) - if isinstance(parsed, list): - return parsed - elif isinstance(parsed, dict): - return [parsed] - - except json.JSONDecodeError as e: - logger.warning(f"Failed to parse JSON experience response: {e}") - - return [] - - -def get_trajectory_context(trajectory: Trajectory, step_sequence: list[Message]) -> str: - """Get context of step sequence within trajectory""" - try: - # Find position of step sequence in trajectory - start_idx = 0 - for i, step in enumerate(trajectory.messages): - if step == step_sequence[0]: - start_idx = i - break - - # Extract before and after context - context_before = trajectory.messages[max(0, start_idx - 2) : start_idx] - context_after = trajectory.messages[start_idx + len(step_sequence) : start_idx + len(step_sequence) + 2] - - context = f"Query: {trajectory.metadata.get('query', 'N/A')}\n" - - if context_before: - context += ( - "Previous steps:\n" - + "\n".join( - [f"- {step.content[:100]}..." for step in context_before], - ) - + "\n" - ) - - if context_after: - context += "Following steps:\n" + "\n".join([f"- {step.content[:100]}..." for step in context_after]) - - return context - - except Exception as e: - logger.error(f"Error getting trajectory context: {e}") - return f"Query: {trajectory.metadata.get('query', 'N/A')}" - - -def extract_content(text: str, language_tag: str = "json", greedy: bool = False): - """Extracts content from Markdown code blocks and parses it if the tag is JSON.""" - quantifier = ".*" if greedy else ".*?" - pattern = rf"```\s*{re.escape(language_tag)}\s*({quantifier})\s*```" - match = re.search(pattern, text, re.DOTALL) - - if not match: - return None - - content = match.group(1).strip() - - if language_tag == "json": - try: - return json.loads(content) - except json.JSONDecodeError: - return None - else: - return content - - -def deduplicate_memories(memories: list[MemoryNode]) -> list[MemoryNode]: - """Deduplicates a list of memories by memory ID.""" - seen_memories: dict[str, MemoryNode] = {} - for memory in memories: - if memory.memory_id not in seen_memories: - seen_memories[memory.memory_id] = memory - return list(seen_memories.values()) diff --git a/reme/core/utils/logger_utils.py b/reme/core/utils/logger_utils.py deleted file mode 100644 index cc141724..00000000 --- a/reme/core/utils/logger_utils.py +++ /dev/null @@ -1,59 +0,0 @@ -"""Logging configuration module for application-wide tracing.""" - -import os -import sys -from datetime import datetime - - -def init_logger( - log_dir: str = "logs", - level: str = "INFO", - log_to_console: bool = True, - log_to_file: bool = True, -) -> None: - """Initialize the logger with both file and console handlers. - - Args: - log_dir: Directory path for log files - level: Logging level (DEBUG, INFO, WARNING, ERROR, CRITICAL) - log_to_console: Whether to print logs to console/screen - log_to_file: Whether to persist logs to files under log_dir - """ - from loguru import logger - - # Remove default handler to avoid duplicate logs - logger.remove() - - # Configure colorized standard output logging if enabled - if log_to_console: - logger.add( - sink=sys.stdout, - level=level, - format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {file}:{line} | {function} | {message}", - colorize=True, - ) - - # Try to configure file-based logging (skip if permission denied) - if log_to_file: - try: - # Ensure the logging directory exists - os.makedirs(log_dir, exist_ok=True) - - # Generate filename based on the current timestamp - # Use dashes instead of colons for Windows compatibility - current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") - log_filename = f"{current_ts}.log" - log_filepath = os.path.join(log_dir, log_filename) - - # Configure file-based logging with rotation and compression - logger.add( - log_filepath, - level=level, - rotation="00:00", - retention="7 days", - compression="zip", - encoding="utf-8", - format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {file}:{line} | {function} | {message}", - ) - except Exception as e: - logger.error(f"Error configuring file logging: {e}") diff --git a/reme/core/utils/logo_utils.py b/reme/core/utils/logo_utils.py deleted file mode 100644 index d85a2073..00000000 --- a/reme/core/utils/logo_utils.py +++ /dev/null @@ -1,85 +0,0 @@ -"""Terminal branding and configuration display utilities.""" - -import importlib.metadata -from typing import TYPE_CHECKING - -from rich.console import Console, Group -from rich.panel import Panel -from rich.table import Table -from rich.text import Text - -if TYPE_CHECKING: - from ..schema import ServiceConfig - - -def get_version(package_name: str) -> str: - """Return the installed version of a package or 'unknown'.""" - try: - return importlib.metadata.version(package_name) - except importlib.metadata.PackageNotFoundError: - return "" - - -def print_logo(service_config: "ServiceConfig"): - """Print a stylized ASCII logo and service metadata to the console.""" - ascii_art = [ - r" ██████╗ ███████╗ ███╗ ███╗ ███████╗ ", - r" ██╔══██╗ ██╔════╝ ████╗ ████║ ██╔════╝ ", - r" ██████╔╝ █████╗ ██╔████╔██║ █████╗ ", - r" ██╔══██╗ ██╔══╝ ██║╚██╔╝██║ ██╔══╝ ", - r" ██║ ██║ ███████╗ ██║ ╚═╝ ██║ ███████╗ ", - r" ╚═╝ ╚═╝ ╚══════╝ ╚═╝ ╚═╝ ╚══════╝ ", - ] - - start_color = (85, 239, 196) - end_color = (162, 155, 254) - - logo_text = Text() - for line in ascii_art: - line_len = max(1, len(line) - 1) - for i, char in enumerate(line): - # Calculate gradient shift per character - ratio = i / line_len - rgb = tuple(int(s + (e - s) * ratio) for s, e in zip(start_color, end_color)) - logo_text.append(char, style=f"bold rgb({rgb[0]},{rgb[1]},{rgb[2]})") - logo_text.append("\n") - - # Layout configuration info - info_table = Table.grid(padding=(0, 1)) - info_table.add_column(style="bold", justify="center") - info_table.add_column(style="bold cyan", justify="left") - info_table.add_column(style="white", justify="left") - - # Add core service info - info_table.add_row("📦", "Backend:", service_config.backend) - - match service_config.backend: - case "http": - host, port = service_config.http.host, service_config.http.port - info_table.add_row("🔗", "URL:", f"http://{host}:{port}") - info_table.add_row("📚", "FastAPI:", Text(get_version("fastapi"), style="dim")) - case "mcp": - mcp = service_config.mcp - transport = mcp.transport if mcp.transport else "stdio" - info_table.add_row("🚌", "Transport:", transport) - if transport != "stdio": - url = f"http://{mcp.host}:{mcp.port}" - if transport == "sse": - url += "/sse" - info_table.add_row("🔗", "URL:", url) - info_table.add_row("📚", "FastMCP:", Text(get_version("fastmcp"), style="dim")) - - info_table.add_row("🚀", "ReMe:", Text(get_version("reme-ai"), style="dim")) - - # Render layout within a panel - panel = Panel( - Group(logo_text, info_table), - title=service_config.app_name, - title_align="left", - border_style="dim", - padding=(1, 4), - expand=False, - ) - - # use justify="center" to adjust position - Console().print(Group("\n", panel, "\n")) diff --git a/reme/core/utils/mcp_client.py b/reme/core/utils/mcp_client.py deleted file mode 100644 index 4c13be9f..00000000 --- a/reme/core/utils/mcp_client.py +++ /dev/null @@ -1,121 +0,0 @@ -"""Module for managing Model Context Protocol (MCP) server connections.""" - -import os -import re -from contextlib import asynccontextmanager -from typing import Any - -from mcp import ClientSession, StdioServerParameters, Tool -from mcp.client.sse import sse_client -from mcp.client.stdio import stdio_client -from mcp.client.streamable_http import streamablehttp_client -from mcp.types import CallToolResult, TextContent - -from ..schema import ToolCall - - -class MCPClient: - """A client manager for handling multiple MCP transport protocols.""" - - def __init__(self, config: dict): - """Initialize the client with server configuration.""" - self.config = config - - @staticmethod - def _infer_transport_type(cfg: dict[str, Any]) -> str: - """Infer the transport type based on configuration fields.""" - if "command" in cfg: - return "stdio" - - if "url" in cfg: - url = cfg["url"].lower() - if url.endswith("/sse") or "sse" in url: - return "sse" - return "streamable-http" - - raise ValueError(f"Could not infer transport type for: {cfg}") - - def _replace_env_vars(self, data: str | dict | list) -> Any: - """Replace environment variable placeholders in configuration.""" - if isinstance(data, str): - return re.sub(r"\$\{(\w+)\}", lambda m: os.getenv(m.group(1), m.group(0)), data) - if isinstance(data, dict): - return {k: self._replace_env_vars(v) for k, v in data.items()} - if isinstance(data, list): - return [self._replace_env_vars(i) for i in data] - return data - - @asynccontextmanager - async def _get_transport(self, cfg: dict[str, Any]): - """Context manager to yield the appropriate MCP transport.""" - # Pop 'type' if present, otherwise infer it - t_type = cfg.pop("type", None) or self._infer_transport_type(cfg) - - try: - if t_type == "stdio": - params = StdioServerParameters( - command=cfg["command"], - args=cfg.get("args", []), - env=cfg.get("env", None), - ) - async with stdio_client(params) as transport: - yield transport - elif t_type == "sse": - async with sse_client(**cfg) as transport: - yield transport - elif t_type == "streamable-http": - async with streamablehttp_client(**cfg) as transport: - yield transport - else: - raise NotImplementedError(f"Unsupported transport: {t_type}") - finally: - pass # Ensure proper cleanup - - @asynccontextmanager - async def connect_to_server(self, server_name: str): - """Establish a session with the specified MCP server.""" - server_config = self.config.get("mcpServers", {}).get(server_name) - if not server_config: - raise ValueError(f"Config for '{server_name}' not found.") - - # Process environment variables and transport selection - cfg = self._replace_env_vars(server_config) - - async with self._get_transport(cfg) as (read, write): - async with ClientSession(read, write) as session: - await session.initialize() - yield session - - async def list_tools(self, server_name: str) -> list[Tool]: - """Retrieve available tools from a specific server.""" - async with self.connect_to_server(server_name) as session: - result = await session.list_tools() - return result.tools - - async def list_tool_calls(self, server_name: str, return_dict: bool = True) -> list[dict | ToolCall]: - """Retrieve available tools from a specific server.""" - tools = await self.list_tools(server_name) - tool_calls: list[ToolCall] = [ToolCall.from_mcp_tool(tool) for tool in tools] - if return_dict: - return [tool_call.simple_input_dump() for tool_call in tool_calls] - - return tool_calls - - async def call_tool( - self, - server_name: str, - tool_name: str, - arguments: dict[str, Any], - parse_text_result: bool = False, - ) -> CallToolResult | str: - """Execute a tool on a specific server.""" - async with self.connect_to_server(server_name) as session: - tool_results: CallToolResult = await session.call_tool(tool_name, arguments) - if not parse_text_result: - return tool_results - - text_result = [] - for block in tool_results.content: - if isinstance(block, TextContent): - text_result.append(block.text) - return "\n".join(text_result) diff --git a/reme/core/utils/pydantic_config_parser.py b/reme/core/utils/pydantic_config_parser.py deleted file mode 100644 index 1f45e4ca..00000000 --- a/reme/core/utils/pydantic_config_parser.py +++ /dev/null @@ -1,198 +0,0 @@ -"""Parser for Pydantic config models with YAML and CLI argument support.""" - -import inspect -import json -from pathlib import Path -from typing import Any, TypeVar - -import yaml -from loguru import logger -from pydantic import BaseModel - -T = TypeVar("T", bound=BaseModel) - - -class PydanticConfigParser: - """Parser that loads and merges Pydantic configs from YAML files and CLI args.""" - - def __init__(self, config_class: type[T], default_config: str = "default"): - """Initialize parser with a Pydantic config class. - - Args: - config_class: Pydantic BaseModel class to validate configs against. - default_config: Default config file name to use if not specified in args. - """ - self.config_class = config_class - self.default_config = default_config - self.config_dict: dict = {} - - def _deep_merge(self, base_dict: dict, update_dict: dict) -> dict: - """Recursively merge two dictionaries.""" - result = base_dict.copy() - for key, value in update_dict.items(): - if key in result and isinstance(result[key], dict) and isinstance(value, dict): - result[key] = self._deep_merge(result[key], value) - else: - result[key] = value - return result - - @staticmethod - def _convert_value(value_str: str) -> Any: - """Convert string value to appropriate Python type.""" - value_str = value_str.strip() - lower_str = value_str.lower() - - # Boolean and None conversion - if lower_str in ("true", "false"): - return lower_str == "true" - if lower_str in ("none", "null"): - return None - - # Numeric conversion - if "e" in lower_str or "." in value_str: - try: - return float(value_str) - except ValueError: - pass - else: - try: - return int(value_str) - except ValueError: - pass - - # JSON conversion for complex types - try: - return json.loads(value_str) - except (json.JSONDecodeError, ValueError): - return value_str - - @staticmethod - def load_from_yaml(yaml_path: str | Path) -> dict: - """Load configuration from YAML file. - - Args: - yaml_path: Path to YAML configuration file. - - Returns: - Dictionary containing configuration data. - - Raises: - FileNotFoundError: If YAML file does not exist. - """ - if isinstance(yaml_path, str): - yaml_path = Path(yaml_path) - - if not yaml_path.exists(): - raise FileNotFoundError(f"Configuration file does not exist: {yaml_path}") - - with yaml_path.open(encoding="utf-8") as f: - return yaml.safe_load(f) or {} - - def merge_configs(self, *config_dicts: dict) -> dict: - """Merge multiple config dictionaries in order. - - Args: - *config_dicts: Variable number of config dictionaries to merge. - - Returns: - Merged configuration dictionary. - """ - result = {} - for config_dict in config_dicts: - result = self._deep_merge(result, config_dict) - return result - - def parse_dot_notation(self, dot_list: list[str]) -> dict: - """Parse dot notation strings into nested dictionary. - - Args: - dot_list: List of strings in format "key.subkey=value". - - Returns: - Nested dictionary representation of dot notation. - """ - config_dict = {} - for item in dot_list: - if "=" not in item: - continue - - key_path, value_str = item.split("=", 1) - keys = key_path.split(".") - - # Build nested dictionary - current = config_dict - for key in keys[:-1]: - current = current.setdefault(key, {}) - current[keys[-1]] = self._convert_value(value_str) - - return config_dict - - def _find_config_path(self, config_name: str) -> Path: - """Find config file path, trying parser directory first then current directory.""" - if not config_name.endswith(".yaml"): - config_name += ".yaml" - - # Try parser class directory first - config_path = Path(inspect.getfile(self.__class__)).parent / config_name - if config_path.exists(): - logger.info(f"load config={config_path}") - return config_path - - # Try current directory - logger.warning(f"config={config_path} not found, try {config_name}") - config_path = Path(config_name) - if not config_path.exists(): - raise FileNotFoundError(f"config={config_path} not found") - return config_path - - def parse_args(self, *args: str, **kwargs) -> T: - """Parse CLI arguments and load configs from YAML files.""" - configs_to_merge = [self.config_class().model_dump()] - - # Separate config file path from other arguments - config = "" - filter_args = [] - for arg in args: - if "=" not in arg: - continue - arg = arg.lstrip("-") - if arg.startswith(("c=", "config=")): - config = arg.split("=", 1)[1] - else: - filter_args.append(arg) - - # Use default config if not specified - config = config or self.default_config - - # Load each config file - for single_config in (c.strip() for c in config.split(",") if c.strip()): - config_path = self._find_config_path(single_config) - configs_to_merge.append(self.load_from_yaml(config_path)) - - # Apply CLI overrides - if filter_args: - configs_to_merge.append(self.parse_dot_notation(filter_args)) - - if kwargs: - configs_to_merge.append(kwargs) - - # Merge all configs and validate - self.config_dict = self.merge_configs(*configs_to_merge) - return self.config_class.model_validate(self.config_dict, extra="allow") - - def update_config(self, **kwargs) -> T: - """Update current config with new values using kwargs. - - Args: - **kwargs: Key-value pairs where __ in keys represents nested levels. - - Returns: - Updated and validated Pydantic config instance. - """ - # Convert kwargs to dot notation and parse - dot_list = [f"{key.replace('__', '.')}={value}" for key, value in kwargs.items()] - override_config = self.parse_dot_notation(dot_list) - - # Merge with existing config - final_config = self.merge_configs(self.config_dict, override_config) - return self.config_class.model_validate(final_config, extra="allow") diff --git a/reme/core/utils/pydantic_utils.py b/reme/core/utils/pydantic_utils.py deleted file mode 100644 index 06b1f4fe..00000000 --- a/reme/core/utils/pydantic_utils.py +++ /dev/null @@ -1,66 +0,0 @@ -""" -Utility module for dynamic Pydantic model generation based on schema definitions. -""" - -from typing import Any, Literal - -from pydantic import create_model, Field - -from . import snake_to_camel -from ..enumeration import JsonSchemaEnum -from ..schema import ToolAttr, Request - -TYPE_MAPPING = {str(t): t.value for t in JsonSchemaEnum} - - -def create_pydantic_model(name: str, parameters: ToolAttr | None = None) -> type[Request]: - """ - Recursively generates a Pydantic model from a ToolAttr schema definition. - """ - fields = {} - - if not parameters or not parameters.properties: - return create_model(f"{snake_to_camel(name)}Model", __base__=Request) - - for field_name, attr in parameters.properties.items(): - # 1. Determine the base field type - if attr.type == "object" and attr.properties: - # Handle nested objects recursively - field_type = create_pydantic_model(field_name, attr) - - elif attr.type == "array" and attr.items: - # Handle array/list types - if isinstance(attr.items, ToolAttr): - if attr.items.type == "object": - inner_type = create_pydantic_model(f"{field_name}_item", attr.items) - else: - inner_type = TYPE_MAPPING.get(attr.items.type, Any) - field_type = list[inner_type] - else: - # Fallback for simple dictionary item definitions - field_type = list[Any] - - else: - # Handle primitive types - field_type = TYPE_MAPPING.get(attr.type, Any) - - # 2. Handle enumeration constraints - if attr.enum: - # Dynamically create a Literal type from the enum list - field_type = Literal[tuple(attr.enum)] # type: ignore - - # 3. Determine requirement status and default values - is_required = False - if parameters.required and field_name in parameters.required: - is_required = True - - # 4. Construct Field metadata - field_info = Field(default=... if is_required else None, description=attr.description) - - if not is_required: - field_type = field_type | None - - fields[field_name] = (field_type, field_info) - - # Dynamically construct the final Pydantic model class - return create_model(f"{snake_to_camel(name)}Model", **fields, __base__=Request) diff --git a/reme/core/utils/pyseekdb_conn.py b/reme/core/utils/pyseekdb_conn.py deleted file mode 100644 index f1c76e2c..00000000 --- a/reme/core/utils/pyseekdb_conn.py +++ /dev/null @@ -1,57 +0,0 @@ -"""Build ``pyseekdb.Client`` / ``AdminClient`` kwargs for embedded vs remote OceanBase / seekdb.""" - -from __future__ import annotations - -DEFAULT_SEEKDB_PORT = 2881 -DEFAULT_SEEKDB_USER = "root" -DEFAULT_SEEKDB_DATABASE = "test" - - -def parse_host_port(host: str | None, port: int | None) -> tuple[str | None, int | None]: - """Resolve remote ``host`` / ``port`` (``port`` defaults to :data:`DEFAULT_SEEKDB_PORT`).""" - if host and host.strip(): - return host.strip(), port if port is not None else DEFAULT_SEEKDB_PORT - return None, None - - -def build_pyseekdb_client_kwargs( - *, - path: str | None = None, - database: str, - host: str | None = None, - port: int | None = None, - user: str | None = None, - password: str = "", -) -> tuple[bool, dict]: - """Return ``(is_remote, kwargs)`` for ``pyseekdb.Client``. - - Remote when ``host`` is set. Embedded: optional ``path`` for the data directory; if - omitted, pyseekdb uses its default (typically ``seekdb.db`` under the CWD). - """ - h, p = parse_host_port(host, port) - if h: - return True, { - "host": h, - "port": p, - "database": database, - "user": user if user is not None else DEFAULT_SEEKDB_USER, - "password": password, - } - kw: dict = {"database": database} - if path: - kw["path"] = path - return False, kw - - -def admin_kwargs_from_client_kwargs(client_kw: dict) -> dict: - """Strip ``database`` for ``AdminClient`` (admin uses system DB).""" - if "path" in client_kw: - return {"path": client_kw["path"]} - if "host" in client_kw: - return { - "host": client_kw["host"], - "port": client_kw["port"], - "user": client_kw["user"], - "password": client_kw["password"], - } - return {} diff --git a/reme/core/utils/singleton.py b/reme/core/utils/singleton.py deleted file mode 100644 index 8c2c071b..00000000 --- a/reme/core/utils/singleton.py +++ /dev/null @@ -1,21 +0,0 @@ -"""Module providing a decorator to implement the Singleton design pattern.""" - -import threading - - -def singleton(cls): - """A class decorator that ensures only one instance of a class exists.""" - - # Dictionary to cache the single instance of the class - _instance = {} - _lock = threading.Lock() - - def _singleton(*args, **kwargs): - """Return the existing instance or create a new one if it doesn't exist.""" - with _lock: - if cls not in _instance: - # Create and store the instance if it's the first call - _instance[cls] = cls(*args, **kwargs) - return _instance[cls] - - return _singleton diff --git a/reme/core/utils/std_logger.py b/reme/core/utils/std_logger.py deleted file mode 100644 index bfa6b32a..00000000 --- a/reme/core/utils/std_logger.py +++ /dev/null @@ -1,122 +0,0 @@ -"""Standard logging module configuration with loguru-like features.""" - -import logging -import os -import sys -from datetime import datetime -from logging.handlers import TimedRotatingFileHandler - -# Store created logger instances -_loggers: dict[str, logging.Logger] = {} - - -class CustomFormatter(logging.Formatter): - """Custom formatter with colorized output support.""" - - # ANSI color codes - COLORS = { - logging.DEBUG: "\033[36m", # Cyan - logging.INFO: "\033[32m", # Green - logging.WARNING: "\033[33m", # Yellow - logging.ERROR: "\033[31m", # Red - logging.CRITICAL: "\033[35m", # Magenta - } - RESET = "\033[0m" - - def __init__(self, fmt: str, colorize: bool = False): - super().__init__(fmt) - self.colorize = colorize - - def format(self, record: logging.LogRecord) -> str: - # Add custom attribute: simplified filename and line number - record.file_line = f"{record.filename}:{record.lineno}" - - if self.colorize: - color = self.COLORS.get(record.levelno, self.RESET) - record.levelname = f"{color}{record.levelname}{self.RESET}" - - return super().format(record) - - -def get_loggerv2( - name: str = "reme", - log_dir: str = "logs", - level: str = "INFO", - log_to_console: bool = True, - log_to_file: bool = True, - log_file_prefix: str = "reme", - rotation: str = "midnight", - retention_days: int = 7, - force_update: bool = False, -) -> logging.Logger: - """Get a configured logger instance. - - Args: - name: Logger name for distinguishing different loggers. - log_dir: Directory path for log files. - level: Logging level (DEBUG, INFO, WARNING, ERROR, CRITICAL). - log_to_console: Whether to output logs to console. - log_to_file: Whether to output logs to file. - log_file_prefix: Prefix for log file names (e.g., 'reme' -> 'reme_2024-01-01.log'). - rotation: Log rotation time, defaults to midnight. - retention_days: Number of days to retain log files. - force_update: Whether to force update the logger configuration even if it already exists. - - Returns: - Configured Logger instance. - """ - # Return existing logger if already created and not force updating - if name in _loggers and not force_update: - return _loggers[name] - - # Create new logger without using root logger - logger = logging.getLogger(name) - logger.setLevel(getattr(logging, level.upper(), logging.INFO)) - logger.propagate = False # Do not propagate to root logger - - # Clear existing handlers - logger.handlers.clear() - - # Log format - log_format = "%(asctime)s | %(levelname)s | %(file_line)s | %(funcName)s | %(message)s" - - # Configure file logging - if log_to_file: - try: - os.makedirs(log_dir, exist_ok=True) - current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") - log_filename = f"{log_file_prefix}_{current_ts}.log" - log_filepath = os.path.join(log_dir, log_filename) - - file_handler = TimedRotatingFileHandler( - log_filepath, - when=rotation, - interval=1, - backupCount=retention_days, - encoding="utf-8", - ) - file_handler.setLevel(getattr(logging, level.upper(), logging.INFO)) - file_handler.setFormatter(CustomFormatter(log_format, colorize=False)) - file_handler.suffix = "%Y-%m-%d" - logger.addHandler(file_handler) - - except Exception as e: - logger.error(f"Error configuring file logging: {e}") - - # Configure console logging - if log_to_console: - console_handler = logging.StreamHandler(sys.stdout) - console_handler.setLevel(getattr(logging, level.upper(), logging.INFO)) - console_handler.setFormatter(CustomFormatter(log_format, colorize=True)) - logger.addHandler(console_handler) - - # Cache logger - _loggers[name] = logger - return logger - - -def get_logger(): - """Get a configured logger instance using loguru.""" - from loguru import logger - - return logger diff --git a/reme/core/utils/time.py b/reme/core/utils/time.py deleted file mode 100644 index d99db735..00000000 --- a/reme/core/utils/time.py +++ /dev/null @@ -1,93 +0,0 @@ -""" -Utility module for timing function execution with log metadata preservation. -""" - -import datetime -import functools -import inspect -import time -from typing import Any, Callable, TypeVar, cast - -from loguru import logger - -# Type variable to preserve the signature of the decorated callable -F = TypeVar("F", bound=Callable[..., Any]) - - -def get_now_time() -> str: - """Get current timestamp in YYYY-MM-DD HH:MM:SS format. - - Returns: - str: Current timestamp string in format 'YYYY-MM-DD HH:MM:SS'. - """ - return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") - - -def timer(func: F) -> F: - """ - Decorator that logs execution time and patches log records with original function metadata. - """ - # Extract original function metadata to ensure logs point to the correct source - func_name = func.__name__ - try: - # Retrieve the source file path and the starting line number - file_path = inspect.getsourcefile(func) or "unknown" - _, line_no = inspect.getsourcelines(func) - except Exception: - file_path = "unknown" - line_no = 0 - - def create_patcher(instance): - """Creates a patcher with runtime class information.""" - # Get the actual class name at runtime if this is a method - if instance is not None and hasattr(instance, "__class__"): - class_name = instance.__class__.__name__ - display_name = f"{class_name}.{func_name}" - else: - display_name = func_name - - def patcher(record): - """Modifies the log record to reflect the decorated function's location.""" - record["function"] = display_name - record["file"].name = file_path.split("/")[-1] - record["file"].path = file_path - record["line"] = line_no - - return patcher - - @functools.wraps(func) - async def async_wrapper(*args: Any, **kwargs: Any) -> Any: - """Timer wrapper for asynchronous functions.""" - start_time = time.perf_counter() - try: - return await func(*args, **kwargs) - finally: - duration = time.perf_counter() - start_time - # Get the instance (self) if this is a method call - instance = args[0] if args else None - patcher = create_patcher(instance) - # Use patch to inject metadata instead of relying on stack depth - logger.patch(patcher).info( - "========== cost={:.6f}s ==========", - duration, - ) - - @functools.wraps(func) - def sync_wrapper(*args: Any, **kwargs: Any) -> Any: - """Timer wrapper for synchronous functions.""" - start_time = time.perf_counter() - try: - return func(*args, **kwargs) - finally: - duration = time.perf_counter() - start_time - # Get the instance (self) if this is a method call - instance = args[0] if args else None - patcher = create_patcher(instance) - logger.patch(patcher).info( - "========== cost={:.6f}s ==========", - duration, - ) - - if inspect.iscoroutinefunction(func): - return cast(F, async_wrapper) - return cast(F, sync_wrapper) diff --git a/reme/core/vector_store/__init__.py b/reme/core/vector_store/__init__.py deleted file mode 100644 index 4c01c748..00000000 --- a/reme/core/vector_store/__init__.py +++ /dev/null @@ -1,51 +0,0 @@ -"""vector store""" - -from .base_vector_store import BaseVectorStore -from .chroma_vector_store import ChromaVectorStore -from .es_vector_store import ESVectorStore -from .hologres_store import HologresVectorStore -from .local_vector_store import LocalVectorStore -from .pgvector_store import PGVectorStore -from .qdrant_vector_store import QdrantVectorStore -from ..registry_factory import R - -__all__ = [ - "BaseVectorStore", - "ChromaVectorStore", - "ESVectorStore", - "HologresVectorStore", - "LocalVectorStore", - "PGVectorStore", - "QdrantVectorStore", -] - -R.vector_stores.register("chroma")(ChromaVectorStore) -R.vector_stores.register("es")(ESVectorStore) -R.vector_stores.register("hologres")(HologresVectorStore) -R.vector_stores.register("local")(LocalVectorStore) -R.vector_stores.register("pgvector")(PGVectorStore) -R.vector_stores.register("qdrant")(QdrantVectorStore) - -try: - from .obvec_vector_store import ObVecVectorStore - - R.vector_stores.register("obvec")(ObVecVectorStore) - __all__.append("ObVecVectorStore") -except ImportError: - pass - -try: - from .zvec_vector_store import ZvecVectorStore - - R.vector_stores.register("zvec")(ZvecVectorStore) - __all__.append("ZvecVectorStore") -except ImportError: - pass - -try: - from .seekdb_vector_store import SeekdbVectorStore - - R.vector_stores.register("seekdb")(SeekdbVectorStore) - __all__.append("SeekdbVectorStore") -except ImportError: - pass diff --git a/reme/core/vector_store/base_vector_store.py b/reme/core/vector_store/base_vector_store.py deleted file mode 100644 index 63791fe8..00000000 --- a/reme/core/vector_store/base_vector_store.py +++ /dev/null @@ -1,123 +0,0 @@ -"""Base vector store interface for managing vector embeddings and similarity search.""" - -from abc import ABC, abstractmethod -from pathlib import Path - -from ..embedding import BaseEmbeddingModel -from ..schema import VectorNode - - -class BaseVectorStore(ABC): - """Abstract base class defining the interface for vector storage and retrieval.""" - - def __init__( - self, - collection_name: str, - db_path: str | Path, - embedding_model: BaseEmbeddingModel, - **kwargs, - ): - """Initialize the vector store with a collection name and an embedding model.""" - self.collection_name: str = collection_name - self.db_path: Path = Path(db_path) - self.embedding_model: BaseEmbeddingModel = embedding_model - self.kwargs: dict = kwargs - - async def get_node_embedding(self, node: VectorNode) -> VectorNode: - """Generate and assign embedding for a single vector node.""" - return await self.embedding_model.get_node_embedding(node) - - async def get_node_embeddings(self, nodes: list[VectorNode]) -> list[VectorNode]: - """Generate and assign embeddings for multiple vector nodes.""" - return await self.embedding_model.get_node_embeddings(nodes) - - async def get_embedding(self, query: str) -> list[float]: - """Convert a single text query into vector embedding using the configured model.""" - return await self.embedding_model.get_embedding(query) - - async def get_embeddings(self, queries: list[str]) -> list[list[float]]: - """Convert multiple text queries into vector embeddings using the configured model.""" - return await self.embedding_model.get_embeddings(queries) - - async def reset_collection(self, collection_name: str): - """Change the name of the current collection.""" - self.collection_name = collection_name - await self.create_collection(collection_name) - - @abstractmethod - async def list_collections(self) -> list[str]: - """Retrieve a list of all existing collection names in the store.""" - - @abstractmethod - async def create_collection(self, collection_name: str, **kwargs) -> None: - """Create a new vector collection with the specified name and configuration.""" - - @abstractmethod - async def delete_collection(self, collection_name: str, **kwargs) -> None: - """Permanently remove a collection from the vector store.""" - - @abstractmethod - async def copy_collection(self, collection_name: str, **kwargs) -> None: - """Duplicate the current collection to a new one with the given name.""" - - @abstractmethod - async def insert(self, nodes: VectorNode | list[VectorNode], **kwargs) -> None: - """Add one or more vector nodes into the current collection.""" - - @abstractmethod - async def search(self, query: str, limit: int = 5, filters: dict | None = None, **kwargs) -> list[VectorNode]: - """Find the most similar vector nodes based on a text query.""" - - @abstractmethod - async def delete(self, vector_ids: str | list[str], **kwargs) -> None: - """Remove specific vectors from the collection using their identifiers.""" - - @abstractmethod - async def delete_all(self, **kwargs) -> None: - """Remove all vectors from the collection.""" - - @abstractmethod - async def update(self, nodes: VectorNode | list[VectorNode], **kwargs) -> None: - """Update the data or metadata of existing vectors in the collection.""" - - @abstractmethod - async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode]: - """Fetch specific vector nodes from the collection by their IDs.""" - - async def dump(self) -> list[VectorNode]: - """Dump the vector store to a list of vector nodes.""" - return await self.list() - - async def load(self, nodes: list[VectorNode]): - """Load the vector store from a list of vector nodes.""" - await self.delete_all() - if nodes: - await self.insert(nodes) - - @abstractmethod - async def list( - self, - filters: dict | None = None, - limit: int | None = None, - sort_key: str | None = None, - reverse: bool = True, - ) -> list[VectorNode]: - """Retrieve vectors from the collection that match the given filters. - - Args: - filters: Dictionary of filter conditions to match vectors - limit: Maximum number of vectors to return - sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting - reverse: If True, sort in descending order; if False, sort in ascending order - """ - - async def start(self) -> None: - """Initialize the vector store and ensure the collection exists. - - This method should be called after instantiation to perform async initialization. - Subclasses should call super().start() to ensure collection creation. - """ - await self.create_collection(self.collection_name) - - async def close(self) -> None: - """Release resources and close active connections to the vector store.""" diff --git a/reme/core/vector_store/chroma_vector_store.py b/reme/core/vector_store/chroma_vector_store.py deleted file mode 100644 index c6bfc1c9..00000000 --- a/reme/core/vector_store/chroma_vector_store.py +++ /dev/null @@ -1,444 +0,0 @@ -"""ChromaDB vector store implementation for the ReMe framework.""" - -from pathlib import Path -from typing import Any - -from loguru import logger - -from .base_vector_store import BaseVectorStore -from ..embedding import BaseEmbeddingModel -from ..schema import VectorNode - -_CHROMADB_IMPORT_ERROR: Exception | None = None - -try: - import chromadb - from chromadb.config import Settings -except Exception as e: - _CHROMADB_IMPORT_ERROR = e - chromadb = None - Settings = None - - -class ChromaVectorStore(BaseVectorStore): - """ChromaDB-based vector store implementation for local or remote storage.""" - - def __init__( - self, - collection_name: str, - db_path: str | Path, - embedding_model: BaseEmbeddingModel, - client: chromadb.ClientAPI | None = None, - host: str | None = None, - port: int | None = None, - api_key: str | None = None, - tenant: str | None = None, - database: str | None = None, - **kwargs, - ): - """Initialize the ChromaDB vector store with the provided configuration.""" - if _CHROMADB_IMPORT_ERROR is not None: - raise ImportError( - "ChromaDB requires extra dependencies. Install with `pip install chromadb`", - ) from _CHROMADB_IMPORT_ERROR - - super().__init__( - collection_name=collection_name, - db_path=db_path, - embedding_model=embedding_model, - **kwargs, - ) - - self.client: chromadb.ClientAPI - self.collection: chromadb.Collection - self.is_local = client is None and not (api_key and tenant) and not (host and port) - - if client: - self.client = client - elif api_key and tenant: - logger.info("Initializing ChromaDB Cloud client") - self.client = chromadb.CloudClient( - api_key=api_key, - tenant=tenant, - database=database or "default", - ) - elif host and port: - logger.info(f"Initializing ChromaDB HTTP client at {host}:{port}") - self.client = chromadb.HttpClient(host=host, port=port) - else: - self.client = None # Will be initialized in start() - - self.collection: chromadb.Collection | None = None - - @staticmethod - def _parse_results( - results: dict, - include_score: bool = False, - ) -> list[VectorNode]: - """Convert ChromaDB query results into a list of VectorNode objects.""" - nodes = [] - - ids = results.get("ids", []) - documents = results.get("documents", []) - metadatas = results.get("metadatas", []) - embeddings = results.get("embeddings") if results.get("embeddings") is not None else [] - distances = results.get("distances") if results.get("distances") is not None else [] - - if ids and isinstance(ids[0], list): - ids = ids[0] if ids else [] - documents = documents[0] if documents else [] - metadatas = metadatas[0] if metadatas else [] - embeddings = embeddings[0] if embeddings and len(embeddings) > 0 else [] - distances = distances[0] if distances and len(distances) > 0 else [] - - for i, vector_id in enumerate(ids): - metadata = metadatas[i] if i < len(metadatas) and metadatas[i] else {} - - if include_score and distances and i < len(distances): - metadata["score"] = 1.0 - distances[i] - - node = VectorNode( - vector_id=vector_id, - content=documents[i] if i < len(documents) and documents[i] else "", - vector=embeddings[i] if len(embeddings) > i else None, - metadata=metadata, - ) - nodes.append(node) - - return nodes - - @staticmethod - def _generate_where_clause(filters: dict | None) -> dict | None: - """Convert the universal filter format to a ChromaDB-compatible where clause. - - Supports two filter formats: - 1. Range query: {"field": [start_value, end_value]} - filters for field >= start_value AND field <= end_value - 2. Exact match: {"field": value} - filters for field == value - """ - if not filters: - return None - - def convert_condition(k: str, v: Any) -> dict | list | None: - """Convert a single filter condition to ChromaDB operator format. - - Returns: - - dict for simple conditions - - list of dicts for range queries (which need to be wrapped in $and) - - None for wildcard filters - """ - if v == "*": - return None - # New syntax: [start, end] represents a range query - if isinstance(v, list) and len(v) == 2: - # Range query: field >= v[0] AND field <= v[1] - # ChromaDB requires separate conditions combined with $and - return [ - {k: {"$gte": v[0]}}, - {k: {"$lte": v[1]}}, - ] - if isinstance(v, dict): - chroma_condition = {} - for op, val in v.items(): - mapping = { - "eq": "$eq", - "ne": "$ne", - "gt": "$gt", - "gte": "$gte", - "lt": "$lt", - "lte": "$lte", - "in": "$in", - "nin": "$nin", - } - chroma_op = mapping.get(op, "$eq") - chroma_condition[k] = {chroma_op: val} - return chroma_condition - # Exact match for non-list values - return {k: {"$eq": v}} - - processed_filters = [] - - for key, value in filters.items(): - if key == "$or": - or_conditions = [] - for condition in value: - or_condition = {} - for sub_key, sub_value in condition.items(): - converted = convert_condition(sub_key, sub_value) - if converted: - if isinstance(converted, list): - # Range query in OR condition - need to wrap in $and - or_conditions.append({"$and": converted}) - else: - or_condition.update(converted) - if or_condition: - or_conditions.append(or_condition) - if len(or_conditions) > 1: - processed_filters.append({"$or": or_conditions}) - elif len(or_conditions) == 1: - processed_filters.append(or_conditions[0]) - - elif key == "$and": - for condition in value: - for sub_key, sub_value in condition.items(): - converted = convert_condition(sub_key, sub_value) - if converted: - if isinstance(converted, list): - # Range query - add each condition separately - processed_filters.extend(converted) - else: - processed_filters.append(converted) - elif key == "$not": - continue - else: - converted = convert_condition(key, value) - if converted: - if isinstance(converted, list): - # Range query - add each condition separately - processed_filters.extend(converted) - else: - processed_filters.append(converted) - - if not processed_filters: - return None - return processed_filters[0] if len(processed_filters) == 1 else {"$and": processed_filters} - - async def list_collections(self) -> list[str]: - """Retrieve a list of all existing collection names.""" - return [col.name for col in self.client.list_collections()] - - async def create_collection(self, collection_name: str, **kwargs): - """Create a new collection with specified distance metrics and metadata.""" - distance_metric = kwargs.get("distance_metric", "cosine") - metadata = kwargs.get("metadata", {}) - metadata["hnsw:space"] = distance_metric - new_collection = self.client.get_or_create_collection(name=collection_name, metadata=metadata) - if collection_name == self.collection_name: - self.collection = new_collection - logger.info(f"Created collection `{collection_name}`") - - async def delete_collection(self, collection_name: str, **kwargs): - """Delete a specified collection from the database.""" - try: - self.client.delete_collection(name=collection_name) - deleted = True - except Exception as _e: - logger.warning(f"Failed to delete collection {collection_name}: {_e}") - deleted = False - if deleted and collection_name == self.collection_name: - self.collection = None - logger.info(f"Deleted collection {collection_name}") - - async def copy_collection(self, collection_name: str, **kwargs): - """Copy all data from the current collection to a new collection.""" - source_data = self.collection.get(include=["documents", "metadatas", "embeddings"]) - if not source_data["ids"]: - logger.warning(f"Source collection {self.collection_name} is empty") - return - - target_collection = self.client.get_or_create_collection( - name=collection_name, - metadata={"hnsw:space": "cosine"}, - ) - target_collection.add( - ids=source_data["ids"], - documents=source_data["documents"], - metadatas=source_data["metadatas"], - embeddings=source_data["embeddings"], - ) - logger.info(f"Copied collection {self.collection_name} to {collection_name}") - - async def insert(self, nodes: VectorNode | list[VectorNode], **kwargs): - """Insert vector nodes into the current collection in batches.""" - if isinstance(nodes, VectorNode): - nodes = [nodes] - if not nodes: - return - - # Batch generate embeddings for nodes that need them - nodes_without_vectors = [node for node in nodes if node.vector is None] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - # Create a mapping for quick lookup - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_insert = [vector_map.get(n.vector_id, n) if n.vector is None else n for n in nodes] - else: - nodes_to_insert = nodes - - batch_size = kwargs.get("batch_size", 100) - - for i in range(0, len(nodes_to_insert), batch_size): - batch_nodes = nodes_to_insert[i : i + batch_size] - self.collection.add( - ids=[n.vector_id for n in batch_nodes], - documents=[n.content for n in batch_nodes], - embeddings=[n.vector for n in batch_nodes], - metadatas=[n.metadata for n in batch_nodes], - ) - logger.info(f"Inserted {len(nodes_to_insert)} nodes into {self.collection_name}") - - async def search( - self, - query: str, - limit: int = 5, - filters: dict | None = None, - **kwargs, - ) -> list[VectorNode]: - """Search for the most similar vector nodes based on a text query.""" - query_vector = await self.get_embedding(query) - where_clause = self._generate_where_clause(filters) - include_embeddings = kwargs.get("include_embeddings", False) - - include: list = ["documents", "metadatas", "distances"] - if include_embeddings: - include.append("embeddings") - results = self.collection.query( - query_embeddings=[query_vector], - n_results=limit, - where=where_clause, - include=include, - ) - nodes = self._parse_results(results, include_score=True) - - score_threshold = kwargs.get("score_threshold") - if score_threshold is not None: - nodes = [n for n in nodes if n.metadata.get("score", 0) >= score_threshold] - return nodes - - async def delete(self, vector_ids: str | list[str], **kwargs): - """Delete specific vector nodes by their IDs.""" - if isinstance(vector_ids, str): - vector_ids = [vector_ids] - if not vector_ids: - return - - self.collection.delete(ids=vector_ids) - logger.info(f"Deleted {len(vector_ids)} nodes from {self.collection_name}") - - async def delete_all(self, **kwargs): - """Remove all vectors from the collection.""" - # Get all IDs in the collection - result = self.collection.get() - count = 0 - if result and result.get("ids"): - ids = result["ids"] - if ids: - self.collection.delete(ids=ids) - count = len(ids) - logger.info(f"Deleted all {count} nodes from {self.collection_name}") - - async def update(self, nodes: VectorNode | list[VectorNode], **kwargs): - """Update existing vector nodes with new content or metadata.""" - if isinstance(nodes, VectorNode): - nodes = [nodes] - if not nodes: - return - - # Batch generate embeddings for nodes that need them - nodes_without_vectors = [node for node in nodes if node.vector is None and node.content] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - # Create a mapping for quick lookup - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_update = [vector_map.get(n.vector_id, n) if n.vector is None and n.content else n for n in nodes] - else: - nodes_to_update = nodes - - self.collection.upsert( - ids=[n.vector_id for n in nodes_to_update], - documents=[n.content for n in nodes_to_update], - embeddings=[n.vector for n in nodes_to_update], - metadatas=[n.metadata for n in nodes_to_update], - ) - logger.info(f"Updated {len(nodes_to_update)} nodes in {self.collection_name}") - - async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode] | None: - """Fetch vector nodes by their IDs from the collection.""" - is_single = isinstance(vector_ids, str) - ids = [vector_ids] if is_single else vector_ids - - results = self.collection.get(ids=ids, include=["documents", "metadatas", "embeddings"]) - nodes = self._parse_results(results) - return nodes[0] if is_single and nodes else (nodes if not is_single else None) - - async def list( - self, - filters: dict | None = None, - limit: int | None = None, - sort_key: str | None = None, - reverse: bool = False, - ) -> list[VectorNode]: - """List vector nodes matching optional metadata filters. - - Args: - filters: Dictionary of filter conditions to match vectors - limit: Maximum number of vectors to return - sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting - reverse: If True, sort in descending order; if False, sort in ascending order - """ - where_clause = self._generate_where_clause(filters) - - # If sorting is needed, fetch all records first, then apply limit after sorting - fetch_limit = None if sort_key else limit - - results = self.collection.get( - where=where_clause, - limit=fetch_limit, - include=["documents", "metadatas", "embeddings"], - ) - nodes = self._parse_results(results) - - # Apply sorting if sort_key is provided - if sort_key: - # Sort with proper handling of None and missing values - def sort_key_func(node): - value = node.metadata.get(sort_key) - if value is None: - # Return appropriate default based on reverse flag - return float("-inf") if not reverse else float("inf") - return value - - nodes.sort(key=sort_key_func, reverse=reverse) - - # Apply limit after sorting - if limit is not None: - nodes = nodes[:limit] - - return nodes - - async def count(self) -> int: - """Return the total number of vectors in the current collection.""" - return self.collection.count() - - async def reset(self): - """Reset the current collection by clearing all its data.""" - logger.warning(f"Resetting collection {self.collection_name}...") - await self.delete_collection(self.collection_name) - - self.collection = self.client.get_or_create_collection( - name=self.collection_name, - metadata={"hnsw:space": "cosine"}, - ) - logger.info(f"Collection {self.collection_name} has been reset") - - async def start(self) -> None: - """Initialize the ChromaDB collection. - - Creates or retrieves the collection with cosine similarity metric. - For local mode, creates the db_path directory if it doesn't exist. - """ - if self.is_local: - self.db_path.mkdir(parents=True, exist_ok=True) - logger.info(f"Initializing local ChromaDB at {self.db_path}") - self.client = chromadb.PersistentClient( - path=str(self.db_path), - settings=Settings(anonymized_telemetry=False), - ) - self.collection = self.client.get_or_create_collection( - name=self.collection_name, - metadata={"hnsw:space": "cosine"}, - ) - logger.info(f"ChromaDB collection {self.collection_name} initialized") - - async def close(self): - """Close the vector store and log the shutdown process.""" - logger.info(f"ChromaDB vector store for collection {self.collection_name} closed") diff --git a/reme/core/vector_store/es_vector_store.py b/reme/core/vector_store/es_vector_store.py deleted file mode 100644 index 332bfa68..00000000 --- a/reme/core/vector_store/es_vector_store.py +++ /dev/null @@ -1,535 +0,0 @@ -"""Elasticsearch vector store implementation for ReMe. - -This module provides an Elasticsearch-based vector store that implements the BaseVectorStore -interface for high-performance dense vector storage and retrieval. -""" - -from pathlib import Path -from typing import Any - -from loguru import logger - -from .base_vector_store import BaseVectorStore -from ..embedding import BaseEmbeddingModel -from ..schema import VectorNode - -_ELASTICSEARCH_IMPORT_ERROR: Exception | None = None - -try: - from elasticsearch import AsyncElasticsearch - from elasticsearch.helpers import async_bulk -except Exception as e: - _ELASTICSEARCH_IMPORT_ERROR = e - AsyncElasticsearch = None - async_bulk = None - - -class ESVectorStore(BaseVectorStore): - """Elasticsearch-based vector store for dense vector storage and kNN search.""" - - def __init__( - self, - collection_name: str, - db_path: str | Path, - embedding_model: BaseEmbeddingModel, - hosts: str | list[str] | None = None, - basic_auth: tuple[str, str] | None = None, - cloud_id: str | None = None, - api_key: str | None = None, - verify_certs: bool = True, - headers: dict[str, str] | None = None, - **kwargs, - ): - """Initialize the Elasticsearch client and vector store configuration. - - Args: - collection_name: Name of the Elasticsearch index (converted to lowercase). - db_path: Database path (not used for remote Elasticsearch, kept for API consistency). - embedding_model: Model instance used to generate vector embeddings. - hosts: Connection host(s) for the Elasticsearch cluster. - basic_auth: Credentials for basic authentication. - cloud_id: Deployment ID for Elastic Cloud. - api_key: API key for authentication. - verify_certs: Enable or disable SSL certificate verification. - headers: Custom HTTP headers for requests. - **kwargs: Additional configuration passed to the base class. - """ - if _ELASTICSEARCH_IMPORT_ERROR is not None: - raise ImportError( - "Elasticsearch requires extra dependencies. Install with `pip install elasticsearch`", - ) from _ELASTICSEARCH_IMPORT_ERROR - - # Elasticsearch requires lowercase index names - collection_name = collection_name.lower() - - super().__init__( - collection_name=collection_name, - db_path=db_path, - embedding_model=embedding_model, - **kwargs, - ) - - # Initialize AsyncElasticsearch client - self.client = AsyncElasticsearch( - hosts=hosts, - cloud_id=cloud_id, - api_key=api_key, - basic_auth=basic_auth, - verify_certs=verify_certs, - headers=headers or {}, - ) - - async def list_collections(self) -> list[str]: - """List all available index names in the Elasticsearch cluster.""" - aliases = await self.client.indices.get_alias() - return list(aliases.keys()) - - async def create_collection(self, collection_name: str, **kwargs): - """Create a new index with dense vector mappings for kNN search. - - Args: - collection_name: Name of the index to create. - **kwargs: Settings like dimensions, similarity, shards, and replicas. - """ - collection_name = collection_name.lower() - - if await self.client.indices.exists(index=collection_name): - return - - dimensions = kwargs.get("dimensions", self.embedding_model.dimensions) - similarity = kwargs.get("similarity", "cosine") - number_of_shards = kwargs.get("number_of_shards", 5) - number_of_replicas = kwargs.get("number_of_replicas", 1) - refresh_interval = kwargs.get("refresh_interval", "1s") - - index_settings = { - "settings": { - "index": { - "number_of_replicas": number_of_replicas, - "number_of_shards": number_of_shards, - "refresh_interval": refresh_interval, - }, - }, - "mappings": { - "properties": { - "vector_id": {"type": "keyword"}, - "content": {"type": "text"}, - "vector": { - "type": "dense_vector", - "dims": dimensions, - "index": True, - "similarity": similarity, - }, - "metadata": {"type": "object", "enabled": True}, - }, - }, - } - - if not await self.client.indices.exists(index=collection_name): - await self.client.indices.create(index=collection_name, body=index_settings) - logger.info(f"Created index {collection_name} with dimensions={dimensions}") - else: - logger.info(f"Index {collection_name} already exists") - - async def delete_collection(self, collection_name: str, **kwargs): - """Permanently delete an Elasticsearch index. - - Args: - collection_name: Name of the index to delete. - **kwargs: Additional parameters for the deletion request. - """ - collection_name = collection_name.lower() - - if await self.client.indices.exists(index=collection_name): - await self.client.indices.delete(index=collection_name) - logger.info(f"Deleted index {collection_name}") - else: - logger.warning(f"Index {collection_name} does not exist") - - async def copy_collection(self, collection_name: str, **kwargs): - """Reindex the current collection into a new index with identical mappings. - - Args: - collection_name: Name of the destination index. - **kwargs: Additional parameters for the reindexing process. - """ - collection_name = collection_name.lower() - - current_index = await self.client.indices.get(index=self.collection_name) - current_settings = current_index[self.collection_name] - - settings_to_copy = current_settings.get("settings", {}).copy() - if "index" in settings_to_copy: - index_settings = settings_to_copy["index"].copy() - internal_keys = [ - "uuid", - "creation_date", - "provided_name", - "version", - "store", - "routing", - "replication", - ] - for key in internal_keys: - index_settings.pop(key, None) - settings_to_copy["index"] = index_settings - - await self.client.indices.create( - index=collection_name, - body={ - "settings": settings_to_copy, - "mappings": current_settings.get("mappings", {}), - }, - ) - - await self.client.reindex( - body={ - "source": {"index": self.collection_name}, - "dest": {"index": collection_name}, - }, - ) - - logger.info(f"Copied collection {self.collection_name} to {collection_name}") - - async def insert(self, nodes: VectorNode | list[VectorNode], refresh: bool = True, **kwargs): - """Insert nodes into the index, generating embeddings if missing. - - Args: - nodes: Single or multiple VectorNode objects to index. - refresh: If True, makes the operation visible to search immediately. - **kwargs: Additional insertion options. - """ - if isinstance(nodes, VectorNode): - nodes = [nodes] - - nodes_without_vectors = [node for node in nodes if node.vector is None] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_insert = [vector_map.get(n.vector_id, n) if n.vector is None else n for n in nodes] - else: - nodes_to_insert = nodes - - actions = [] - for node in nodes_to_insert: - action = { - "_index": self.collection_name, - "_id": node.vector_id, - "_source": { - "vector_id": node.vector_id, - "content": node.content, - "vector": node.vector, - "metadata": node.metadata, - }, - } - actions.append(action) - - success, failed = await async_bulk(self.client, actions, raise_on_error=False) - - if failed: - logger.warning(f"Failed to insert {len(failed)} documents") - - logger.info(f"Inserted {success} documents into {self.collection_name}") - - if refresh: - await self.client.indices.refresh(index=self.collection_name) - - async def search( - self, - query: str, - limit: int = 5, - filters: dict | None = None, - **kwargs, - ) -> list[VectorNode]: - """Perform a kNN similarity search based on a text query. - - Args: - query: The text to search for. - limit: Maximum number of nearest neighbors to return. - filters: Metadata filters for exact match or 'IN' operations. - **kwargs: Search parameters like num_candidates or score_threshold. - - Returns: - List of VectorNode objects ordered by similarity. - """ - query_vector = await self.get_embedding(query) - num_candidates = kwargs.get("num_candidates", limit * 2) - - search_query: dict = { - "knn": { - "field": "vector", - "query_vector": query_vector, - "k": limit, - "num_candidates": num_candidates, - }, - "size": limit, - } - - if filters: - filter_conditions = [] - for key, value in filters.items(): - # New syntax: [start, end] represents a range query - if isinstance(value, list) and len(value) == 2: - # Range query: field >= value[0] AND field <= value[1] - filter_conditions.append( - { - "range": { - f"metadata.{key}": { - "gte": value[0], - "lte": value[1], - }, - }, - }, - ) - else: - # Exact match - filter_conditions.append({"term": {f"metadata.{key}": value}}) - search_query["knn"]["filter"] = {"bool": {"must": filter_conditions}} - - response = await self.client.search(index=self.collection_name, body=search_query) - - results = [] - for hit in response["hits"]["hits"]: - source = hit["_source"] - node = VectorNode( - vector_id=source.get("vector_id", hit["_id"]), - content=source.get("content", ""), - vector=source.get("vector"), - metadata=source.get("metadata", {}), - ) - node.metadata["score"] = hit["_score"] - results.append(node) - - return results - - async def delete(self, vector_ids: str | list[str], refresh: bool = True, **kwargs): - """Delete specific vectors from the index by their IDs. - - Args: - vector_ids: Single ID or list of IDs to remove. - refresh: If True, refreshes the index after deletion. - **kwargs: Additional deletion parameters. - """ - if isinstance(vector_ids, str): - vector_ids = [vector_ids] - - actions = [] - for vector_id in vector_ids: - actions.append( - { - "_op_type": "delete", - "_index": self.collection_name, - "_id": vector_id, - }, - ) - - success, failed = await async_bulk( - self.client, - actions, - raise_on_error=False, - raise_on_exception=False, - ) - - if failed: - logger.warning(f"Failed to delete {len(failed)} documents") - - logger.info(f"Deleted {success} documents from {self.collection_name}") - - if refresh: - await self.client.indices.refresh(index=self.collection_name) - - async def delete_all(self, **kwargs): - """Remove all vectors from the collection. - - Args: - **kwargs: Additional deletion parameters. - """ - response = await self.client.delete_by_query( - index=self.collection_name, - body={"query": {"match_all": {}}}, - ) - - deleted_count = response.get("deleted", 0) - logger.info(f"Deleted all {deleted_count} documents from {self.collection_name}") - - refresh = kwargs.get("refresh", True) - if refresh: - await self.client.indices.refresh(index=self.collection_name) - - async def update(self, nodes: VectorNode | list[VectorNode], refresh: bool = True, **kwargs): - """Update existing documents with new content or metadata. - - Args: - nodes: Single or multiple VectorNode objects with updated data. - refresh: If True, refreshes the index after update. - **kwargs: Additional update parameters. - """ - if isinstance(nodes, VectorNode): - nodes = [nodes] - - nodes_without_vectors = [node for node in nodes if node.vector is None and node.content] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_update = [vector_map.get(n.vector_id, n) if n.vector is None and n.content else n for n in nodes] - else: - nodes_to_update = nodes - - actions = [] - for node in nodes_to_update: - doc: dict = { - "vector_id": node.vector_id, - "content": node.content, - "metadata": node.metadata, - } - if node.vector is not None: - doc["vector"] = node.vector - - actions.append( - { - "_op_type": "update", - "_index": self.collection_name, - "_id": node.vector_id, - "doc": doc, - }, - ) - - success, failed = await async_bulk( - self.client, - actions, - raise_on_error=False, - raise_on_exception=False, - ) - - if failed: - logger.warning(f"Failed to update {len(failed)} documents") - - logger.info(f"Updated {success} documents in {self.collection_name}") - - if refresh: - await self.client.indices.refresh(index=self.collection_name) - - async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode]: - """Fetch documents by their IDs from the current index. - - Args: - vector_ids: Single ID or list of IDs to retrieve. - - Returns: - A single VectorNode or a list of VectorNodes. - """ - single_result = isinstance(vector_ids, str) - if single_result: - vector_ids = [vector_ids] - - response = await self.client.mget( - index=self.collection_name, - body={"ids": vector_ids}, - ) - - results = [] - for doc in response["docs"]: - if doc.get("found"): - source = doc["_source"] - node = VectorNode( - vector_id=source.get("vector_id", doc["_id"]), - content=source.get("content", ""), - vector=source.get("vector"), - metadata=source.get("metadata", {}), - ) - results.append(node) - else: - logger.warning(f"Document with ID {doc['_id']} not found") - - return results[0] if single_result and results else results - - async def list( - self, - filters: dict | None = None, - limit: int | None = None, - sort_key: str | None = None, - reverse: bool = False, - ) -> list[VectorNode]: - """Retrieve a list of nodes filtered by metadata or limit. - - Args: - filters: Optional metadata filtering criteria. - limit: Maximum number of nodes to return. - sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting - reverse: If True, sort in descending order; if False, sort in ascending order - - Returns: - A list of matching VectorNode objects. - """ - query: dict[str, Any] = {"query": {"match_all": {}}} - - if filters: - filter_conditions = [] - for key, value in filters.items(): - # New syntax: [start, end] represents a range query - if isinstance(value, list) and len(value) == 2: - # Range query: field >= value[0] AND field <= value[1] - filter_conditions.append( - { - "range": { - f"metadata.{key}": { - "gte": value[0], - "lte": value[1], - }, - }, - }, - ) - else: - # Exact match - filter_conditions.append({"term": {f"metadata.{key}": value}}) - query["query"] = {"bool": {"must": filter_conditions}} - - # Add sorting to the Elasticsearch query if sort_key is provided - if sort_key: - query["sort"] = [ - { - f"metadata.{sort_key}": { - "order": "desc" if reverse else "asc", - }, - }, - ] - - if limit: - query["size"] = limit - else: - query["size"] = 10000 - - response = await self.client.search(index=self.collection_name, body=query) - - results = [] - for hit in response["hits"]["hits"]: - source = hit["_source"] - node = VectorNode( - vector_id=source.get("vector_id", hit["_id"]), - content=source.get("content", ""), - vector=source.get("vector"), - metadata=source.get("metadata", {}), - ) - results.append(node) - - return results - - async def reset_collection(self, collection_name: str): - """Reset collection with lowercase conversion for Elasticsearch compatibility.""" - collection_name = collection_name.lower() - self.collection_name = collection_name - await self.create_collection(collection_name) - logger.info(f"Collection reset to {collection_name}") - - async def start(self) -> None: - """Initialize the Elasticsearch index. - - Creates the index with dense vector mappings if it doesn't exist. - """ - await super().start() - logger.info(f"Elasticsearch index {self.collection_name} initialized") - - async def close(self): - """Terminate the Elasticsearch client session and release resources.""" - await self.client.close() - logger.info("Elasticsearch client connection closed") diff --git a/reme/core/vector_store/hologres_store.py b/reme/core/vector_store/hologres_store.py deleted file mode 100644 index 270dc1bf..00000000 --- a/reme/core/vector_store/hologres_store.py +++ /dev/null @@ -1,633 +0,0 @@ -"""Hologres implementation for vector storage and retrieval.""" - -import json -import re -from pathlib import Path -from typing import Any - -from loguru import logger - -from .base_vector_store import BaseVectorStore -from ..embedding import BaseEmbeddingModel -from ..schema import VectorNode - -_ASYNCPG_IMPORT_ERROR: Exception | None = None - -try: - import asyncpg - from asyncpg import Pool -except Exception as e: - _ASYNCPG_IMPORT_ERROR = e - asyncpg = None - Pool = None - - -class HologresVectorStore(BaseVectorStore): - """Vector store implementation using Hologres for efficient similarity search. - - Hologres uses native float4[] arrays for vector storage with built-in - HGraph index for approximate nearest neighbor search, unlike pgvector - which requires an extension. - """ - - @staticmethod - def _validate_table_name(name: str) -> None: - """Validate table name to prevent SQL injection.""" - if not name: - raise ValueError("Table name cannot be empty") - if len(name) > 63: - raise ValueError(f"Table name too long: {len(name)} characters (max 63)") - if not re.match(r"^[a-zA-Z_][a-zA-Z0-9_]*$", name): - raise ValueError( - f"Invalid table name: {name}. Must start with letter or underscore, " - "and contain only alphanumeric characters and underscores.", - ) - - def __init__( - self, - collection_name: str, - db_path: str | Path, - embedding_model: BaseEmbeddingModel, - host: str = "localhost", - port: int = 80, - database: str = "postgres", - user: str = "postgres", - password: str = "", - schema: str = "public", - min_size: int = 1, - max_size: int = 10, - dsn: str | None = None, - distance_method: str = "Cosine", - **kwargs, - ): - """Initialize the Hologres vector store with connection parameters. - - Args: - collection_name: Name of the collection (table). - db_path: Database path (used by base class). - embedding_model: Embedding model for generating vectors. - host: Hologres host address. - port: Hologres port (default 80 for Hologres). - database: Database name. - user: Database user. - password: Database password. - schema: PostgreSQL schema name (default "public"). - min_size: Minimum connections in pool. - max_size: Maximum connections in pool. - dsn: Full DSN connection string (overrides individual params). - distance_method: Distance method for HGraph index (Cosine, InnerProduct, Euclidean). - """ - if _ASYNCPG_IMPORT_ERROR is not None: - raise ImportError( - "Hologres vector store requires asyncpg. Install with `pip install asyncpg`", - ) from _ASYNCPG_IMPORT_ERROR - - self._validate_table_name(collection_name) - self._validate_table_name(schema) - - super().__init__( - collection_name=collection_name, - db_path=db_path, - embedding_model=embedding_model, - **kwargs, - ) - - self.dsn = dsn - self.host = host - self.port = port - self.database = database - self.user = user - self.password = password - self.schema = schema - self.min_size = min_size - self.max_size = max_size - self.distance_method = distance_method - self._pool: Pool | None = None - self.embedding_model_dims = embedding_model.dimensions - - @property - def _qualified_name(self) -> str: - """Return the schema-qualified table name (e.g. 'my_schema.my_table').""" - return f"{self.schema}.{self.collection_name}" - - def _qualify(self, table_name: str) -> str: - """Return a schema-qualified name for an arbitrary table.""" - return f"{self.schema}.{table_name}" - - @staticmethod - async def _hologres_reset(conn): - """Custom reset for Hologres connections.""" - await conn.execute( - """ - SELECT pg_advisory_unlock_all(); - CLOSE ALL; - RESET ALL; - """, - ) - - async def _get_pool(self) -> Pool: - """Create or return the existing asyncpg connection pool.""" - if self._pool is None: - if self.dsn: - self._pool = await asyncpg.create_pool( - dsn=self.dsn, - min_size=self.min_size, - max_size=self.max_size, - reset=self._hologres_reset, - ) - else: - self._pool = await asyncpg.create_pool( - host=self.host, - port=self.port, - database=self.database, - user=self.user, - password=self.password, - min_size=self.min_size, - max_size=self.max_size, - reset=self._hologres_reset, - ) - - # Ensure schema exists - async with self._pool.acquire() as conn: - await conn.execute(f"CREATE SCHEMA IF NOT EXISTS {self.schema}") - - logger.info(f"Hologres connection pool created for database {self.database}") - - return self._pool - - @staticmethod - def _vector_to_pg_array(vector: list[float]) -> str: - """Convert a Python list of floats to PostgreSQL array literal format.""" - return "{" + ",".join(map(str, vector)) + "}" - - @staticmethod - def _pg_array_to_vector(pg_array) -> list[float] | None: - """Convert a PostgreSQL array result to a Python list of floats.""" - if pg_array is None: - return None - if isinstance(pg_array, list): - return [float(x) for x in pg_array] - # Handle string format like {1.0,2.0,3.0} - raw = str(pg_array) - if raw.startswith("{") and raw.endswith("}"): - return [float(x) for x in raw[1:-1].split(",")] - return None - - async def list_collections(self) -> list[str]: - """List all available table names in the current schema.""" - pool = await self._get_pool() - async with pool.acquire() as conn: - rows = await conn.fetch( - "SELECT table_name FROM information_schema.tables WHERE table_schema = $1", - self.schema, - ) - return [row["table_name"] for row in rows] - - async def create_collection(self, collection_name: str, **kwargs): - """Create a new Hologres table with vector support and HGraph index.""" - self._validate_table_name(collection_name) - pool = await self._get_pool() - dimensions = kwargs.get("dimensions", self.embedding_model_dims) - qualified = self._qualify(collection_name) - - async with pool.acquire() as conn: - create_sql = f""" - CREATE TABLE IF NOT EXISTS {qualified} ( - id TEXT PRIMARY KEY, - content TEXT, - vector float4[] CHECK (array_ndims(vector) = 1 AND array_length(vector, 1) = {dimensions}), - metadata JSONB - ) - WITH ( - vectors = '{{ - "vector": {{ - "algorithm": "HGraph", - "distance_method": "{self.distance_method}", - "builder_params": {{ - "base_quantization_type": "rabitq", - "rabitq_use_fht":true, - "graph_storage_type": "compressed", - "max_total_size_to_merge_mb": 4096, - "max_degree": 64, - "ef_construction": 400, - "precise_quantization_type": "fp32", - "use_reorder": true - }} - }} - }}' - ) - """ - await conn.execute(create_sql) - - logger.info(f"Created Hologres collection {qualified} with dimensions={dimensions}") - - async def delete_collection(self, collection_name: str, **kwargs): - """Remove the specified collection table from the database.""" - self._validate_table_name(collection_name) - pool = await self._get_pool() - qualified = self._qualify(collection_name) - async with pool.acquire() as conn: - await conn.execute(f"DROP TABLE IF EXISTS {qualified}") - logger.info(f"Deleted collection {qualified}") - - async def copy_collection(self, collection_name: str, **kwargs): - """Duplicate the structure and content of the current collection to a new table.""" - self._validate_table_name(collection_name) - pool = await self._get_pool() - qualified_src = self._qualified_name - qualified_dst = self._qualify(collection_name) - - async with pool.acquire() as conn: - columns = await conn.fetch( - """ - SELECT column_name, data_type, udt_name - FROM information_schema.columns - WHERE table_name = $1 AND table_schema = $2 - """, - self.collection_name, - self.schema, - ) - - if not columns: - raise ValueError(f"Source collection {qualified_src} does not exist") - - # Create new table with primary key, then add data - await conn.execute( - f""" - SET hg_experimental_enable_create_table_like_properties = true; - CALL hg_create_table_like('{qualified_dst}', 'select * from {qualified_src}') - """, - ) - await conn.execute(f"INSERT INTO {qualified_dst} SELECT * FROM {qualified_src} ;") - - logger.info(f"Copied collection {qualified_src} to {qualified_dst}") - - async def insert(self, nodes: VectorNode | list[VectorNode], **kwargs): - """Insert or upsert vector nodes into the Hologres collection.""" - if isinstance(nodes, VectorNode): - nodes = [nodes] - - if not nodes: - return - - nodes_without_vectors = [node for node in nodes if node.vector is None] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_insert = [vector_map.get(n.vector_id, n) if n.vector is None else n for n in nodes] - else: - nodes_to_insert = nodes - - pool = await self._get_pool() - data = [ - ( - node.vector_id, - node.content, - node.vector, - json.dumps(node.metadata), - ) - for node in nodes_to_insert - ] - - async with pool.acquire() as conn: - on_conflict = kwargs.get("on_conflict", "update") - - if on_conflict == "update": - await conn.executemany( - f""" - INSERT INTO {self._qualified_name} (id, content, vector, metadata) - VALUES ($1, $2, $3::float4[], $4::jsonb) - ON CONFLICT (id) DO UPDATE SET - content = EXCLUDED.content, - vector = EXCLUDED.vector, - metadata = EXCLUDED.metadata - """, - data, - ) - elif on_conflict == "ignore": - await conn.executemany( - f""" - INSERT INTO {self._qualified_name} (id, content, vector, metadata) - VALUES ($1, $2, $3::float4[], $4::jsonb) - ON CONFLICT (id) DO NOTHING - """, - data, - ) - else: - await conn.executemany( - f""" - INSERT INTO {self._qualified_name} (id, content, vector, metadata) - VALUES ($1, $2, $3::float4[], $4::jsonb) - """, - data, - ) - - logger.info(f"Inserted {len(nodes_to_insert)} documents into {self._qualified_name}") - - @staticmethod - def _build_filter_clause(filters: dict | None) -> tuple[str, list]: - """Generate an SQL WHERE clause and parameter list from a filter dictionary. - - Supports two filter formats: - 1. Range query: {"field": [start_value, end_value]} - 2. Exact match: {"field": value} - """ - if not filters: - return "", [] - - conditions = [] - params = [] - param_idx = 1 - - for key, value in filters.items(): - if not key.replace("_", "").replace(".", "").isalnum(): - raise ValueError( - f"Invalid metadata key: {key}. Only alphanumeric characters, underscore and dot are allowed.", - ) - - if isinstance(value, list) and len(value) == 2: - if isinstance(value[0], (int, float)) and isinstance(value[1], (int, float)): - conditions.append( - f"(metadata->>'{key}')::numeric >= ${param_idx} AND " - f"(metadata->>'{key}')::numeric <= ${param_idx + 1}", - ) - else: - conditions.append(f"metadata->>'{key}' >= ${param_idx} AND metadata->>'{key}' <= ${param_idx + 1}") - params.extend([value[0], value[1]]) - param_idx += 2 - else: - conditions.append(f"metadata->>'{key}' = ${param_idx}") - params.append(str(value)) - param_idx += 1 - - filter_clause = "WHERE " + " AND ".join(conditions) if conditions else "" - return filter_clause, params - - async def search( - self, - query: str, - limit: int = 5, - filters: dict | None = None, - **kwargs, - ) -> list[VectorNode]: - """Perform vector similarity search using Hologres approx_cosine_distance.""" - query_vector = await self.get_embedding(query) - vector_str = self._vector_to_pg_array(query_vector) - pool = await self._get_pool() - - filter_clause, filter_params = self._build_filter_clause(filters) - - # filter_params use $1..$N, limit uses $(N+1) - limit_placeholder = f"${len(filter_params) + 1}" - - async with pool.acquire() as conn: - sql = f""" - SELECT id, content, vector, metadata, - approx_cosine_distance(vector, '{vector_str}') AS distance - FROM {self._qualified_name} - {filter_clause} - ORDER BY distance DESC - LIMIT {limit_placeholder} - """ - rows = await conn.fetch(sql, *filter_params, limit) - - results = [] - score_threshold = kwargs.get("score_threshold") - - for row in rows: - distance = float(row["distance"]) - # approx_cosine_distance returns cosine similarity (higher = more similar) - score = distance - if score_threshold is not None and score < score_threshold: - continue - - vector_data = self._pg_array_to_vector(row["vector"]) - - metadata = row["metadata"] if row["metadata"] else {} - if isinstance(metadata, str): - metadata = json.loads(metadata) - - metadata["score"] = score - metadata["_distance"] = 1 - score - - node = VectorNode( - vector_id=row["id"], - content=row["content"] or "", - vector=vector_data, - metadata=metadata, - ) - results.append(node) - - return results - - async def delete(self, vector_ids: str | list[str], **kwargs): - """Remove specific vector records from the collection by their IDs.""" - if isinstance(vector_ids, str): - vector_ids = [vector_ids] - - if not vector_ids: - return - - pool = await self._get_pool() - async with pool.acquire() as conn: - placeholders = ", ".join([f"${i + 1}" for i in range(len(vector_ids))]) - await conn.execute( - f"DELETE FROM {self._qualified_name} WHERE id IN ({placeholders})", - *vector_ids, - ) - - logger.info(f"Deleted {len(vector_ids)} documents from {self._qualified_name}") - - async def delete_all(self, **kwargs): - """Remove all vectors from the collection.""" - pool = await self._get_pool() - async with pool.acquire() as conn: - result = await conn.execute(f"DELETE FROM {self._qualified_name}") - - logger.info(f"Deleted all documents from {self._qualified_name} result={result}") - - async def update(self, nodes: VectorNode | list[VectorNode], **kwargs): - """Update existing vector nodes with new content, embeddings, or metadata.""" - if isinstance(nodes, VectorNode): - nodes = [nodes] - - if not nodes: - return - - nodes_without_vectors = [node for node in nodes if node.vector is None and node.content] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_update = [vector_map.get(n.vector_id, n) if n.vector is None and n.content else n for n in nodes] - else: - nodes_to_update = nodes - - pool = await self._get_pool() - async with pool.acquire() as conn: - for node in nodes_to_update: - update_fields = [] - params = [] - idx = 1 - - if node.content: - update_fields.append(f"content = ${idx}") - params.append(node.content) - idx += 1 - - if node.vector: - update_fields.append(f"vector = ${idx}::float4[]") - params.append(node.vector) - idx += 1 - - if node.metadata: - update_fields.append(f"metadata = ${idx}::jsonb") - params.append(json.dumps(node.metadata)) - idx += 1 - - if update_fields: - params.append(node.vector_id) - await conn.execute( - f"UPDATE {self._qualified_name} SET {', '.join(update_fields)} WHERE id = ${idx}", - *params, - ) - - logger.info(f"Updated {len(nodes_to_update)} documents in {self._qualified_name}") - - async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode] | None: - """Retrieve vector nodes by their unique identifiers.""" - single_result = isinstance(vector_ids, str) - if single_result: - vector_ids = [vector_ids] - - if not vector_ids: - return [] if not single_result else None - - pool = await self._get_pool() - async with pool.acquire() as conn: - placeholders = ", ".join([f"${i + 1}" for i in range(len(vector_ids))]) - rows = await conn.fetch( - f"SELECT id, content, vector, metadata FROM {self._qualified_name} WHERE id IN ({placeholders})", - *vector_ids, - ) - - results = [] - for row in rows: - vector_data = self._pg_array_to_vector(row["vector"]) - - metadata = row["metadata"] if row["metadata"] else {} - if isinstance(metadata, str): - metadata = json.loads(metadata) - - results.append( - VectorNode( - vector_id=row["id"], - content=row["content"] or "", - vector=vector_data, - metadata=metadata, - ), - ) - - if single_result: - return results[0] if results else None - return results - - async def list( - self, - filters: dict | None = None, - limit: int | None = None, - sort_key: str | None = None, - reverse: bool = False, - ) -> list[VectorNode]: - """Return a list of vector nodes matching the provided filters and limit. - - Args: - filters: Dictionary of filter conditions to match vectors - limit: Maximum number of vectors to return - sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting - reverse: If True, sort in descending order; if False, sort in ascending order - """ - pool = await self._get_pool() - filter_clause, filter_params = self._build_filter_clause(filters) - - order_clause = "" - if sort_key: - order_direction = "DESC" if reverse else "ASC" - order_clause = f"ORDER BY metadata->>'{sort_key}' {order_direction}" - - limit_clause = "" - if limit: - limit_clause = f"LIMIT ${len(filter_params) + 1}" - filter_params.append(limit) - - async with pool.acquire() as conn: - sql = f""" - SELECT id, content, vector, metadata - FROM {self._qualified_name} - {filter_clause} - {order_clause} - {limit_clause} - """ - rows = await conn.fetch(sql, *filter_params) - - results = [] - for row in rows: - vector_data = self._pg_array_to_vector(row["vector"]) - - metadata = row["metadata"] if row["metadata"] else {} - if isinstance(metadata, str): - metadata = json.loads(metadata) - - results.append( - VectorNode( - vector_id=row["id"], - content=row["content"] or "", - vector=vector_data, - metadata=metadata, - ), - ) - - return results - - async def collection_info(self) -> dict[str, Any]: - """Fetch metadata including record count and disk usage for the collection.""" - pool = await self._get_pool() - qualified = self._qualified_name - - async with pool.acquire() as conn: - count = await conn.fetchval(f"SELECT COUNT(*) FROM {qualified}") - size = await conn.fetchval(f"SELECT pg_size_pretty(pg_total_relation_size('{qualified}'))") - - return { - "name": qualified, - "count": count, - "size": size, - } - - async def reset(self): - """Purge all data by dropping and recreating the collection table.""" - logger.warning(f"Resetting collection {self._qualified_name}...") - await self.delete_collection(self.collection_name) - await self.create_collection(self.collection_name) - - async def reset_collection(self, collection_name: str): - """Reset collection with table name validation.""" - self._validate_table_name(collection_name) - self.collection_name = collection_name - await self.create_collection(collection_name) - logger.info(f"Collection reset to {self._qualified_name}") - - async def start(self) -> None: - """Initialize the PGVector store. - - Creates the connection pool and ensures the collection table exists. - """ - await self._get_pool() - await super().start() - logger.info(f"Hologres collection {self._qualified_name} initialized") - - async def close(self): - """Terminate the database connection pool.""" - if self._pool is not None: - await self._pool.close() - self._pool = None - logger.info("Hologres connection pool closed") diff --git a/reme/core/vector_store/local_vector_store.py b/reme/core/vector_store/local_vector_store.py deleted file mode 100644 index fb2a4a1b..00000000 --- a/reme/core/vector_store/local_vector_store.py +++ /dev/null @@ -1,341 +0,0 @@ -"""Local file system vector store implementation for ReMe.""" - -import json -from pathlib import Path - -import numpy as np -from loguru import logger - -from .base_vector_store import BaseVectorStore -from ..embedding import BaseEmbeddingModel -from ..schema import VectorNode -from ..utils import batch_cosine_similarity - - -class LocalVectorStore(BaseVectorStore): - """Local file system-based vector store with in-memory caching. - - All operations are performed in memory after start(). - Changes are persisted to disk on close(). - """ - - def __init__( - self, - collection_name: str, - db_path: str | Path, - embedding_model: BaseEmbeddingModel, - **kwargs, - ): - """Initialize the local vector store with a db_path and collection name.""" - super().__init__( - collection_name=collection_name, - db_path=db_path, - embedding_model=embedding_model, - **kwargs, - ) - # In-memory cache: vector_id -> VectorNode - self._cache: dict[str, VectorNode] = {} - self._dirty: bool = False # Track if cache has unsaved changes - - def _get_collection_path(self, collection_name: str) -> Path: - """Get the file system path for a specific collection.""" - return self.db_path / collection_name - - def _get_node_file_path(self, vector_id: str, collection_name: str | None = None) -> Path: - """Get the JSON file path for a specific vector node.""" - col_path = self._get_collection_path(collection_name or self.collection_name) - return col_path / f"{vector_id}.json" - - def _save_node_to_disk(self, node: VectorNode): - """Save a vector node to a JSON file on disk.""" - file_path = self._get_node_file_path(node.vector_id) - file_path.parent.mkdir(parents=True, exist_ok=True) - with open(file_path, "w", encoding="utf-8") as f: - json.dump(node.model_dump(), f, ensure_ascii=False, indent=2) - - def _load_all_from_disk(self) -> dict[str, VectorNode]: - """Load all vector nodes from disk into a dictionary.""" - col_path = self._get_collection_path(self.collection_name) - if not col_path.exists(): - return {} - - nodes = {} - for file_path in col_path.glob("*.json"): - try: - with open(file_path, "r", encoding="utf-8") as f: - data = json.load(f) - node = VectorNode(**data) - nodes[node.vector_id] = node - except Exception as e: - logger.warning(f"Failed to load node from {file_path}: {e}") - return nodes - - def _flush_to_disk(self): - """Persist all cached nodes to disk.""" - col_path = self._get_collection_path(self.collection_name) - col_path.mkdir(parents=True, exist_ok=True) - - # Remove files that are no longer in cache - existing_files = set(col_path.glob("*.json")) - cached_ids = set(self._cache.keys()) - for file_path in existing_files: - vector_id = file_path.stem - if vector_id not in cached_ids: - file_path.unlink() - - # Write all cached nodes - for node in self._cache.values(): - self._save_node_to_disk(node) - - self._dirty = False - logger.info(f"Flushed {len(self._cache)} nodes to disk") - - @staticmethod - def _match_filters(node: VectorNode, filters: dict | None) -> bool: - """Check if a vector node matches the provided metadata filters. - - Supports two filter formats: - 1. Range query: {"field": [start_value, end_value]} - filters for field >= start_value AND field <= end_value - 2. Exact match: {"field": value} - filters for field == value - """ - if not filters: - return True - - for key, value in filters.items(): - node_value = node.metadata.get(key) - - # New syntax: [start, end] represents a range query - if isinstance(value, list) and len(value) == 2: - # Range query: field >= value[0] AND field <= value[1] - if node_value is None: - return False - try: - # Try numeric comparison - if not value[0] <= node_value <= value[1]: - return False - except TypeError: - # If comparison fails, the filter doesn't match - return False - else: - # Exact match - if node_value != value: - return False - - return True - - async def list_collections(self) -> list[str]: - """List all collection directories in the db_path.""" - if not self.db_path.exists(): - return [] - - return [d.name for d in self.db_path.iterdir() if d.is_dir() and not d.name.startswith(".")] - - async def create_collection(self, collection_name: str, **kwargs): - """Create a new collection directory.""" - col_path = self._get_collection_path(collection_name) - col_path.mkdir(parents=True, exist_ok=True) - logger.info(f"Created collection {collection_name} at {col_path}") - - async def delete_collection(self, collection_name: str, **kwargs): - """Delete a collection directory and all its JSON files.""" - col_path = self._get_collection_path(collection_name) - - if not col_path.exists(): - logger.warning(f"Collection {collection_name} does not exist") - return - - for file_path in col_path.glob("*.json"): - file_path.unlink() - - col_path.rmdir() - logger.info(f"Deleted collection {collection_name}") - - async def copy_collection(self, collection_name: str, **kwargs): - """Copy all nodes from the current collection to a new one.""" - source_path = self._get_collection_path(self.collection_name) - target_path = self._get_collection_path(collection_name) - - if not source_path.exists(): - logger.warning(f"Source collection {self.collection_name} does not exist") - return - - target_path.mkdir(parents=True, exist_ok=True) - - for file_path in source_path.glob("*.json"): - target_file = target_path / file_path.name - target_file.write_text(file_path.read_text(encoding="utf-8"), encoding="utf-8") - - logger.info(f"Copied collection {self.collection_name} to {collection_name}") - - async def insert(self, nodes: VectorNode | list[VectorNode], **kwargs): - """Insert vector nodes into the cache, generating embeddings if necessary.""" - if isinstance(nodes, VectorNode): - nodes = [nodes] - - nodes_without_vectors = [node for node in nodes if node.vector is None] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_insert = [vector_map.get(n.vector_id, n) if n.vector is None else n for n in nodes] - else: - nodes_to_insert = nodes - - for node in nodes_to_insert: - self._cache[node.vector_id] = node - - self._dirty = True - logger.info(f"Inserted {len(nodes_to_insert)} nodes into {self.collection_name}") - - async def search( - self, - query: str, - limit: int = 5, - filters: dict | None = None, - **kwargs, - ) -> list[VectorNode]: - """Search for nodes similar to the query using batch cosine similarity.""" - query_vector = await self.get_embedding(query) - - # Filter nodes from cache - filtered_nodes = [node for node in self._cache.values() if self._match_filters(node, filters)] - - # Separate nodes with and without vectors - nodes_with_vectors = [node for node in filtered_nodes if node.vector is not None] - if not nodes_with_vectors: - return [] - - # Build matrix for batch similarity computation - node_vectors = np.array([node.vector for node in nodes_with_vectors]) - query_matrix = np.array([query_vector]) - - # Compute similarities in batch: shape (1, num_nodes) -> flatten to (num_nodes,) - similarities = batch_cosine_similarity(query_matrix, node_vectors).flatten() - - # Apply score threshold if specified - score_threshold = kwargs.get("score_threshold") - - # Pair nodes with scores and filter/sort - scored_nodes = list(zip(nodes_with_vectors, similarities)) - if score_threshold is not None: - scored_nodes = [(node, score) for node, score in scored_nodes if score >= score_threshold] - - scored_nodes.sort(key=lambda x: x[1], reverse=True) - scored_nodes = scored_nodes[:limit] - - # Attach scores to metadata - results = [] - for node, score in scored_nodes: - node.metadata["score"] = float(score) - results.append(node) - - return results - - async def delete(self, vector_ids: str | list[str], **kwargs): - """Delete specific vector nodes by their IDs from cache.""" - if isinstance(vector_ids, str): - vector_ids = [vector_ids] - - deleted_count = 0 - for vector_id in vector_ids: - if vector_id in self._cache: - del self._cache[vector_id] - deleted_count += 1 - else: - logger.warning(f"Node {vector_id} does not exist") - - if deleted_count > 0: - self._dirty = True - logger.info(f"Deleted {deleted_count} nodes from {self.collection_name}") - - async def delete_all(self, **kwargs): - """Remove all vectors from the cache.""" - count = len(self._cache) - self._cache.clear() - self._dirty = True - logger.info(f"Deleted all {count} nodes from {self.collection_name}") - - async def update(self, nodes: VectorNode | list[VectorNode], **kwargs): - """Update existing vector nodes in the cache.""" - if isinstance(nodes, VectorNode): - nodes = [nodes] - - nodes_without_vectors = [node for node in nodes if node.vector is None and node.content] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_update = [vector_map.get(n.vector_id, n) if n.vector is None and n.content else n for n in nodes] - else: - nodes_to_update = nodes - - updated_count = 0 - for node in nodes_to_update: - if node.vector_id in self._cache: - self._cache[node.vector_id] = node - updated_count += 1 - else: - logger.warning(f"Node {node.vector_id} does not exist, skipping update") - - if updated_count > 0: - self._dirty = True - logger.info(f"Updated {updated_count} nodes in {self.collection_name}") - - async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode]: - """Retrieve one or more vector nodes from cache by their unique IDs.""" - is_single = isinstance(vector_ids, str) - ids = [vector_ids] if is_single else vector_ids - - results = [] - for vector_id in ids: - node = self._cache.get(vector_id) - if node: - results.append(node) - else: - logger.warning(f"Node {vector_id} not found") - - return results[0] if is_single and results else results - - async def list( - self, - filters: dict | None = None, - limit: int | None = None, - sort_key: str | None = None, - reverse: bool = True, - ) -> list[VectorNode]: - """List vector nodes from cache with optional filtering and limits. - - Args: - filters: Dictionary of filter conditions to match vectors - limit: Maximum number of vectors to return - sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting - reverse: If True, sort in descending order; if False, sort in ascending order - """ - filtered_nodes = [node for node in self._cache.values() if self._match_filters(node, filters)] - - # Apply sorting if sort_key is provided - if sort_key: - - def sort_key_func(node): - value = node.metadata.get(sort_key) - if value is None: - return float("-inf") if not reverse else float("inf") - return value - - filtered_nodes.sort(key=sort_key_func, reverse=reverse) - - if limit is not None: - filtered_nodes = filtered_nodes[:limit] - - return filtered_nodes - - async def start(self) -> None: - """Initialize the local vector store and load all nodes into memory.""" - await super().start() - self._cache = self._load_all_from_disk() - self._dirty = False - logger.info(f"Local vector store loaded {len(self._cache)} nodes from {self.collection_name}") - - async def close(self): - """Persist all cached data to disk and close the vector store.""" - if self._dirty: - self._flush_to_disk() - logger.info("Local vector store closed") diff --git a/reme/core/vector_store/obvec_vector_store.py b/reme/core/vector_store/obvec_vector_store.py deleted file mode 100644 index 74231343..00000000 --- a/reme/core/vector_store/obvec_vector_store.py +++ /dev/null @@ -1,461 +0,0 @@ -"""OceanBase / seekdb vector store for ReMe (pyobvector).""" - -from __future__ import annotations - -import json -from pathlib import Path -from typing import Any - -from loguru import logger - -from .base_vector_store import BaseVectorStore -from ..embedding import BaseEmbeddingModel -from ..schema import VectorNode - -_OBVEC_IMPORT_ERROR: Exception | None = None - -try: - from sqlalchemy import Column, JSON, String, text as sa_text - from sqlalchemy.dialects.mysql import LONGTEXT - from pyobvector import IndexParams, ObVecClient, VecIndexType, VECTOR - from pyobvector import cosine_distance, inner_product -except Exception as e: - _OBVEC_IMPORT_ERROR = e - Column = None # type: ignore[misc, assignment] - JSON = None # type: ignore[misc, assignment] - String = None # type: ignore[misc, assignment] - sa_text = None # type: ignore[misc, assignment] - LONGTEXT = None # type: ignore[misc, assignment] - IndexParams = None # type: ignore[misc, assignment] - ObVecClient = None # type: ignore[misc, assignment] - VecIndexType = None # type: ignore[misc, assignment] - VECTOR = None # type: ignore[misc, assignment] - cosine_distance = None # type: ignore[misc, assignment] - inner_product = None # type: ignore[misc, assignment] - -_COL_SELECT = "id, content, vector, metadata" - - -def _is_safe_metadata_key(key: str) -> bool: - return bool(key.replace("_", "").replace(".", "").isalnum()) - - -def _coerce_db_vector(raw: Any) -> list[float] | None: - if raw is None: - return None - if isinstance(raw, list): - return [float(x) for x in raw] - if isinstance(raw, str): - try: - parsed = json.loads(raw) - if isinstance(parsed, list): - return [float(x) for x in parsed] - except (json.JSONDecodeError, TypeError, ValueError): - pass - return None - - -def _coerce_db_metadata(raw: Any) -> dict[str, Any]: - if raw is None: - return {} - if isinstance(raw, dict): - return raw - if isinstance(raw, str): - try: - parsed = json.loads(raw) - if isinstance(parsed, dict): - return parsed - except (json.JSONDecodeError, TypeError): - pass - return {} - - -def _build_metadata_filter_sql(filters: dict[str, Any] | None) -> str: - if not filters: - return "" - - parts: list[str] = [] - for key, value in filters.items(): - if not _is_safe_metadata_key(key): - continue - - path = f"$.{key}" - if isinstance(value, list) and len(value) == 2: - lo, hi = value[0], value[1] - if isinstance(lo, (int, float)) and isinstance(hi, (int, float)): - parts.append( - f"(JSON_EXTRACT(metadata, '{path}') >= {lo} AND " f"JSON_EXTRACT(metadata, '{path}') <= {hi})", - ) - else: - parts.append( - f"(JSON_EXTRACT(metadata, '{path}') >= '{lo}' AND " f"JSON_EXTRACT(metadata, '{path}') <= '{hi}')", - ) - elif isinstance(value, (int, float)): - parts.append(f"JSON_EXTRACT(metadata, '{path}') = {value}") - else: - parts.append(f"JSON_EXTRACT(metadata, '{path}') = '{value}'") - - return " AND ".join(parts) - - -def _format_vector_sql_literal(vector: list[float]) -> str: - return "[" + ",".join(str(float(v)) for v in vector) + "]" - - -def _normalize_embedding_for_ann(raw: Any) -> list[float]: - if hasattr(raw, "tolist"): - raw = raw.tolist() - return [float(x) for x in raw] - - -def _vector_node_from_db_row(row: tuple[Any, ...]) -> VectorNode: - return VectorNode( - vector_id=row[0], - content=row[1] or "", - vector=_coerce_db_vector(row[2]), - metadata=_coerce_db_metadata(row[3]), - ) - - -def _normalize_nodes(nodes: VectorNode | list[VectorNode]) -> list[VectorNode]: - return [nodes] if isinstance(nodes, VectorNode) else list(nodes) - - -def _sql_table(name: str) -> str: - return f"`{name}`" - - -class ObVecVectorStore(BaseVectorStore): - """OceanBase or seekdb vector store for dense vectors and kNN search. - - Args: - index_metric: ``cosine`` or ``ip`` (inner product). Invalid values raise - ``ValueError``; unsupported strings are not mapped to another metric. - """ - - def __init__( - self, - collection_name: str, - db_path: str | Path, - embedding_model: BaseEmbeddingModel, - uri: str = "127.0.0.1:2881", - user: str = "root", - password: str = "", - database: str = "test", - index_metric: str = "cosine", - index_ef_search: int = 100, - **kwargs, - ): - if _OBVEC_IMPORT_ERROR is not None: - raise ImportError( - "ObVecVectorStore requires pyobvector and sqlalchemy. " - "Install with `pip install pyobvector sqlalchemy`", - ) from _OBVEC_IMPORT_ERROR - - super().__init__( - collection_name=collection_name, - db_path=db_path, - embedding_model=embedding_model, - **kwargs, - ) - - key = index_metric.strip().lower() - if key not in ("cosine", "ip"): - raise ValueError( - f"ObVecVectorStore index_metric must be 'cosine' or 'ip', got {index_metric!r}", - ) - self.uri = uri - self.user = user - self.password = password - self.database = database - self.index_metric = key - self.index_ef_search = index_ef_search - - self.client: ObVecClient | None = None - self.embedding_model_dims = embedding_model.dimensions - - def _require_client(self) -> ObVecClient: - if self.client is None: - raise RuntimeError("ObVecVectorStore.start() must be called before this operation") - return self.client - - async def list_collections(self) -> list[str]: - client = self._require_client() - result = client.perform_raw_text_sql(f"SHOW TABLES FROM {_sql_table(self.database)}") - rows = result.fetchall() - return [row[0] for row in rows if row] - - def _table_columns_for_create(self, dimensions: int) -> list[Any]: - return [ - Column("id", String(255), primary_key=True), - Column("content", LONGTEXT), - Column("vector", VECTOR(dimensions)), - Column("metadata", JSON), - ] - - def _hnsw_index_params(self, collection_name: str) -> IndexParams: - metric = "cosine" if self.index_metric == "cosine" else "inner_product" - vidxs = IndexParams() - vidxs.add_index( - "vector", - VecIndexType.HNSW, - f"{collection_name}_vidx", - metric_type=metric, - params={"efSearch": self.index_ef_search}, - ) - return vidxs - - async def create_collection(self, collection_name: str, **kwargs): - client = self._require_client() - dimensions = kwargs.get("dimensions", self.embedding_model_dims) - - if client.check_table_exists(collection_name): - logger.info("Collection {} already exists", collection_name) - return - - columns = self._table_columns_for_create(dimensions) - vidxs = self._hnsw_index_params(collection_name) - - client.create_table_with_index_params( - table_name=collection_name, - columns=columns, - vidxs=vidxs, - ) - logger.info("Created collection {} with dimensions={}", collection_name, dimensions) - - async def delete_collection(self, collection_name: str, **kwargs): - client = self._require_client() - client.drop_table_if_exist(collection_name) - logger.info("Deleted collection {}", collection_name) - - async def copy_collection(self, collection_name: str, **kwargs): - client = self._require_client() - - if not client.check_table_exists(self.collection_name): - raise ValueError(f"Source collection {self.collection_name} does not exist") - - await self.create_collection(collection_name) - - try: - source_data = await self.list(limit=None) - if source_data: - await self.insert(source_data, collection_name=collection_name) - logger.info("Copied collection {} to {}", self.collection_name, collection_name) - except Exception: - try: - client.drop_table_if_exist(collection_name) - except Exception as cleanup_err: - logger.warning("Cleanup after failed copy failed: {}", cleanup_err) - raise - - async def insert(self, nodes: VectorNode | list[VectorNode], **kwargs): - nodes = _normalize_nodes(nodes) - if not nodes: - return - - client = self._require_client() - - need_emb = [n for n in nodes if n.vector is None] - if need_emb: - filled = await self.get_node_embeddings(need_emb) - by_id = {n.vector_id: n for n in filled} - nodes_to_insert = [by_id.get(n.vector_id, n) for n in nodes] - else: - nodes_to_insert = nodes - - data = [ - { - "id": node.vector_id, - "content": node.content, - "vector": node.vector if node.vector is not None else [], - "metadata": node.metadata if node.metadata else {}, - } - for node in nodes_to_insert - ] - target = kwargs.get("collection_name", self.collection_name) - client.insert(table_name=target, data=data) - logger.info("Inserted {} documents into {}", len(nodes_to_insert), target) - - async def search( - self, - query: str, - limit: int = 5, - filters: dict | None = None, - **kwargs, - ) -> list[VectorNode]: - client = self._require_client() - raw_vec = await self.get_embedding(query) - query_vector = _normalize_embedding_for_ann(raw_vec) - dist_fn = cosine_distance if self.index_metric == "cosine" else inner_product - - filter_sql = _build_metadata_filter_sql(filters) - where_parts = [sa_text(filter_sql)] if filter_sql else None - - results = client.ann_search( - table_name=self.collection_name, - vec_data=query_vector, - vec_column_name="vector", - distance_func=dist_fn, - with_dist=True, - topk=limit, - output_column_names=["id", "content", "metadata"], - where_clause=where_parts, - ) - - score_threshold = kwargs.get("score_threshold") - out: list[VectorNode] = [] - for row in results: - if len(row) < 4: - raise RuntimeError( - "ann_search row must have id, content, metadata, distance " f"(got {len(row)} columns)", - ) - vid, content, metadata_raw, distance = row[0], row[1], row[2], row[3] - dist_f = float(distance) - if self.index_metric == "cosine": - score = max(0.0, 1.0 - dist_f / 2.0) - else: - score = max(0.0, dist_f) - if score_threshold is not None and score < score_threshold: - continue - meta = _coerce_db_metadata(metadata_raw) if metadata_raw is not None else {} - meta["score"] = score - meta["_distance"] = dist_f - out.append( - VectorNode( - vector_id=vid, - content=content or "", - vector=None, - metadata=meta, - ), - ) - return out - - async def delete(self, vector_ids: str | list[str], **kwargs): - if isinstance(vector_ids, str): - vector_ids = [vector_ids] - if not vector_ids: - return - client = self._require_client() - client.delete(self.collection_name, ids=vector_ids) - logger.info("Deleted {} documents from {}", len(vector_ids), self.collection_name) - - async def delete_all(self, **kwargs): - client = self._require_client() - client.delete(self.collection_name) - logger.info("Deleted all documents from {}", self.collection_name) - - async def update(self, nodes: VectorNode | list[VectorNode], **kwargs): - nodes = _normalize_nodes(nodes) - if not nodes: - return - - client = self._require_client() - need_emb = [n for n in nodes if n.vector is None and bool(n.content)] - if need_emb: - filled = await self.get_node_embeddings(need_emb) - by_id = {n.vector_id: n for n in filled} - nodes_to_update = [ - by_id.get(n.vector_id, n) if (n.vector is None and bool(n.content)) else n for n in nodes - ] - else: - nodes_to_update = nodes - - for node in nodes_to_update: - updates: list[str] = [] - params: dict[str, Any] = {} - - if node.content is not None: - updates.append("content = :content") - params["content"] = node.content - - if node.vector is not None: - updates.append("vector = :vector") - params["vector"] = _format_vector_sql_literal(node.vector) - - if node.metadata is not None: - updates.append("metadata = :metadata") - params["metadata"] = json.dumps(node.metadata) - - if not updates: - continue - - params["vid"] = node.vector_id - update_sql = f"UPDATE {_sql_table(self.collection_name)} SET {', '.join(updates)} WHERE id = :vid" - with client.engine.connect() as conn: - with conn.begin(): - conn.execute(sa_text(update_sql), params) - - logger.info("Updated {} documents in {}", len(nodes_to_update), self.collection_name) - - async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode] | None: - single = isinstance(vector_ids, str) - if single: - vector_ids = [vector_ids] - if not vector_ids: - return [] if not single else None - - client = self._require_client() - ids_str = "', '".join(vector_ids) - select_sql = f"SELECT {_COL_SELECT} FROM {_sql_table(self.collection_name)} " f"WHERE id IN ('{ids_str}')" - result = client.perform_raw_text_sql(select_sql) - rows = result.fetchall() - parsed = [_vector_node_from_db_row(row) for row in rows if row] - if single: - return parsed[0] if parsed else None - return parsed - - async def list( - self, - filters: dict | None = None, - limit: int | None = None, - sort_key: str | None = None, - reverse: bool = False, - ) -> list[VectorNode]: - client = self._require_client() - select_sql = f"SELECT {_COL_SELECT} FROM {_sql_table(self.collection_name)}" - where_clause = _build_metadata_filter_sql(filters) - if where_clause: - select_sql += f" WHERE {where_clause}" - if sort_key and _is_safe_metadata_key(sort_key): - order = "DESC" if reverse else "ASC" - select_sql += f" ORDER BY JSON_EXTRACT(metadata, '$.{sort_key}') {order}" - if limit is not None: - select_sql += f" LIMIT {limit}" - result = client.perform_raw_text_sql(select_sql) - rows = result.fetchall() - return [_vector_node_from_db_row(row) for row in rows if row] - - async def collection_info(self) -> dict[str, Any]: - """Return collection name and row count.""" - client = self._require_client() - count_sql = f"SELECT COUNT(*) FROM {_sql_table(self.collection_name)}" - result = client.perform_raw_text_sql(count_sql) - row = result.fetchone() - count = row[0] if row else 0 - return {"name": self.collection_name, "count": count} - - async def reset(self): - """Drop and recreate the current collection table.""" - logger.warning("Resetting collection {}...", self.collection_name) - await self.delete_collection(self.collection_name) - await self.create_collection(self.collection_name) - - async def reset_collection(self, collection_name: str): - self.collection_name = collection_name - await self.create_collection(collection_name) - logger.info("Collection reset to {}", collection_name) - - async def start(self) -> None: - self.client = ObVecClient( - uri=self.uri, - user=self.user, - password=self.password, - db_name=self.database, - ) - - await super().start() - logger.info("seekdb / OceanBase vector table {} ready", self.collection_name) - - async def close(self): - self.client = None - logger.info("ObVec client connection closed") diff --git a/reme/core/vector_store/pgvector_store.py b/reme/core/vector_store/pgvector_store.py deleted file mode 100644 index 03805076..00000000 --- a/reme/core/vector_store/pgvector_store.py +++ /dev/null @@ -1,610 +0,0 @@ -"""PostgreSQL pgvector implementation for vector storage and retrieval.""" - -import json -import re -from pathlib import Path -from typing import Any - -from loguru import logger - -from .base_vector_store import BaseVectorStore -from ..embedding import BaseEmbeddingModel -from ..schema import VectorNode - -_ASYNCPG_IMPORT_ERROR: Exception | None = None - -try: - import asyncpg - from asyncpg import Pool -except Exception as e: - _ASYNCPG_IMPORT_ERROR = e - asyncpg = None - Pool = None - - -class PGVectorStore(BaseVectorStore): - """Vector store implementation using PostgreSQL and pgvector for efficient similarity search.""" - - @staticmethod - def _validate_table_name(name: str) -> None: - """Validate table name to prevent SQL injection. - - PostgreSQL table names must: - - Contain only alphanumeric characters and underscores - - Not start with a digit - - Be between 1 and 63 characters - """ - if not name: - raise ValueError("Table name cannot be empty") - if len(name) > 63: - raise ValueError(f"Table name too long: {len(name)} characters (max 63)") - if not re.match(r"^[a-zA-Z_][a-zA-Z0-9_]*$", name): - raise ValueError( - f"Invalid table name: {name}. Must start with letter or underscore, " - "and contain only alphanumeric characters and underscores.", - ) - - def __init__( - self, - collection_name: str, - db_path: str | Path, - embedding_model: BaseEmbeddingModel, - host: str = "localhost", - port: int = 5432, - database: str = "postgres", - user: str = "postgres", - password: str = "", - min_size: int = 1, - max_size: int = 10, - dsn: str | None = None, - use_hnsw: bool = True, - use_diskann: bool = False, - **kwargs, - ): - """Initialize the PGVector store with connection parameters and index settings.""" - if _ASYNCPG_IMPORT_ERROR is not None: - raise ImportError( - "PGVector requires extra dependencies. Install with `pip install asyncpg pgvector`", - ) from _ASYNCPG_IMPORT_ERROR - - # Validate collection name to prevent SQL injection - self._validate_table_name(collection_name) - - super().__init__( - collection_name=collection_name, - db_path=db_path, - embedding_model=embedding_model, - **kwargs, - ) - - self.dsn = dsn - self.host = host - self.port = port - self.database = database - self.user = user - self.password = password - self.min_size = min_size - self.max_size = max_size - self.use_hnsw = use_hnsw - self.use_diskann = use_diskann - self._pool: Pool | None = None - self.embedding_model_dims = embedding_model.dimensions - - async def _get_pool(self) -> Pool: - """Create or return the existing asyncpg connection pool.""" - if self._pool is None: - if self.dsn: - self._pool = await asyncpg.create_pool( - dsn=self.dsn, - min_size=self.min_size, - max_size=self.max_size, - ) - else: - self._pool = await asyncpg.create_pool( - host=self.host, - port=self.port, - database=self.database, - user=self.user, - password=self.password, - min_size=self.min_size, - max_size=self.max_size, - ) - - async with self._pool.acquire() as conn: - await conn.execute("CREATE EXTENSION IF NOT EXISTS vector") - - logger.info(f"PGVector connection pool created for database {self.database}") - - return self._pool - - async def list_collections(self) -> list[str]: - """List all available table names in the current database.""" - pool = await self._get_pool() - async with pool.acquire() as conn: - rows = await conn.fetch( - "SELECT table_name FROM information_schema.tables WHERE table_schema = 'public'", - ) - return [row["table_name"] for row in rows] - - async def create_collection(self, collection_name: str, **kwargs): - """Create a new PostgreSQL table with vector support and appropriate indexing.""" - self._validate_table_name(collection_name) - pool = await self._get_pool() - dimensions = kwargs.get("dimensions", self.embedding_model_dims) - - async with pool.acquire() as conn: - await conn.execute( - f""" - CREATE TABLE IF NOT EXISTS {collection_name} ( - id TEXT PRIMARY KEY, - content TEXT, - vector vector({dimensions}), - metadata JSONB - ) - """, - ) - - if self.use_diskann and dimensions < 2000: - result = await conn.fetchval( - "SELECT 1 FROM pg_extension WHERE extname = 'vectorscale'", - ) - if result: - await conn.execute( - f""" - CREATE INDEX IF NOT EXISTS {collection_name}_diskann_idx - ON {collection_name} - USING diskann (vector) - """, - ) - logger.info(f"Created DiskANN index for collection {collection_name}") - else: - logger.warning("vectorscale extension not available, skipping DiskANN index") - elif self.use_hnsw: - await conn.execute( - f""" - CREATE INDEX IF NOT EXISTS {collection_name}_hnsw_idx - ON {collection_name} - USING hnsw (vector vector_cosine_ops) - """, - ) - logger.info(f"Created HNSW index for collection {collection_name}") - - logger.info(f"Created collection {collection_name} with dimensions={dimensions}") - - async def delete_collection(self, collection_name: str, **kwargs): - """Remove the specified collection table from the database.""" - self._validate_table_name(collection_name) - pool = await self._get_pool() - async with pool.acquire() as conn: - await conn.execute(f"DROP TABLE IF EXISTS {collection_name}") - logger.info(f"Deleted collection {collection_name}") - - async def copy_collection(self, collection_name: str, **kwargs): - """Duplicate the structure and content of the current collection to a new table.""" - self._validate_table_name(collection_name) - pool = await self._get_pool() - - async with pool.acquire() as conn: - columns = await conn.fetch( - """ - SELECT column_name, data_type, udt_name - FROM information_schema.columns - WHERE table_name = $1 AND table_schema = 'public' - """, - self.collection_name, - ) - - if not columns: - raise ValueError(f"Source collection {self.collection_name} does not exist") - - await conn.execute(f"CREATE TABLE {collection_name} AS TABLE {self.collection_name}") - await conn.execute(f"ALTER TABLE {collection_name} ADD PRIMARY KEY (id)") - - if self.use_hnsw: - await conn.execute( - f""" - CREATE INDEX IF NOT EXISTS {collection_name}_hnsw_idx - ON {collection_name} - USING hnsw (vector vector_cosine_ops) - """, - ) - - logger.info(f"Copied collection {self.collection_name} to {collection_name}") - - async def insert(self, nodes: VectorNode | list[VectorNode], **kwargs): - """Insert or upsert vector nodes into the PostgreSQL collection.""" - if isinstance(nodes, VectorNode): - nodes = [nodes] - - if not nodes: - return - - nodes_without_vectors = [node for node in nodes if node.vector is None] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_insert = [vector_map.get(n.vector_id, n) if n.vector is None else n for n in nodes] - else: - nodes_to_insert = nodes - - pool = await self._get_pool() - data = [ - ( - node.vector_id, - node.content, - f"[{','.join(map(str, node.vector))}]", - json.dumps(node.metadata), - ) - for node in nodes_to_insert - ] - - async with pool.acquire() as conn: - on_conflict = kwargs.get("on_conflict", "update") - - if on_conflict == "update": - await conn.executemany( - f""" - INSERT INTO {self.collection_name} (id, content, vector, metadata) - VALUES ($1, $2, $3::vector, $4::jsonb) - ON CONFLICT (id) DO UPDATE SET - content = EXCLUDED.content, - vector = EXCLUDED.vector, - metadata = EXCLUDED.metadata - """, - data, - ) - elif on_conflict == "ignore": - await conn.executemany( - f""" - INSERT INTO {self.collection_name} (id, content, vector, metadata) - VALUES ($1, $2, $3::vector, $4::jsonb) - ON CONFLICT (id) DO NOTHING - """, - data, - ) - else: - await conn.executemany( - f""" - INSERT INTO {self.collection_name} (id, content, vector, metadata) - VALUES ($1, $2, $3::vector, $4::jsonb) - """, - data, - ) - - logger.info(f"Inserted {len(nodes_to_insert)} documents into {self.collection_name}") - - @staticmethod - def _build_filter_clause(filters: dict | None) -> tuple[str, list]: - """Generate an SQL WHERE clause and parameter list from a filter dictionary. - - Supports two filter formats: - 1. Range query: {"field": [start_value, end_value]} - filters for field >= start_value AND field <= end_value - 2. Exact match: {"field": value} - filters for field == value - - Range queries support both numeric and string (e.g., timestamp strings) comparisons. - """ - if not filters: - return "", [] - - conditions = [] - params = [] - param_idx = 1 - - for key, value in filters.items(): - # Sanitize key to prevent SQL injection (only allow alphanumeric and underscore) - if not key.replace("_", "").replace(".", "").isalnum(): - raise ValueError( - f"Invalid metadata key: {key}. Only alphanumeric characters, underscore and dot are allowed.", - ) - - # New syntax: [start, end] represents a range query - if isinstance(value, list) and len(value) == 2: - # Range query: field >= value[0] AND field <= value[1] - # Try numeric comparison first, fall back to text comparison if needed - if isinstance(value[0], (int, float)) and isinstance(value[1], (int, float)): - # Numeric range query - conditions.append( - f"(metadata->>'{key}')::numeric >= ${param_idx} AND " - f"(metadata->>'{key}')::numeric <= ${param_idx + 1}", - ) - else: - # Text range query (works for strings, timestamps, etc.) - conditions.append(f"metadata->>'{key}' >= ${param_idx} AND metadata->>'{key}' <= ${param_idx + 1}") - params.extend([value[0], value[1]]) - param_idx += 2 - else: - # Exact match - conditions.append(f"metadata->>'{key}' = ${param_idx}") - params.append(str(value)) - param_idx += 1 - - filter_clause = "WHERE " + " AND ".join(conditions) if conditions else "" - return filter_clause, params - - async def search( - self, - query: str, - limit: int = 5, - filters: dict | None = None, - **kwargs, - ) -> list[VectorNode]: - """Perform vector similarity search with optional metadata filtering.""" - query_vector = await self.get_embedding(query) - vector_str = f"[{','.join(map(str, query_vector))}]" - pool = await self._get_pool() - - filter_clause, filter_params = self._build_filter_clause(filters) - - # Adjust parameter indices in filter clause to account for $1 being used by vector_str - if filter_clause: - for i in range(len(filter_params), 0, -1): - new_placeholder = f"${i + 1}" - filter_clause = re.sub(rf"\${i}\b", new_placeholder, filter_clause) - - async with pool.acquire() as conn: - sql = f""" - SELECT id, content, vector, metadata, vector <=> $1::vector AS distance - FROM {self.collection_name} - {filter_clause} - ORDER BY distance - LIMIT ${len(filter_params) + 2} - """ - rows = await conn.fetch(sql, vector_str, *filter_params, limit) - - results = [] - score_threshold = kwargs.get("score_threshold") - - for row in rows: - distance = row["distance"] - if score_threshold is not None and distance > score_threshold: - continue - - vector_data = None - if row["vector"]: - vector_str_raw = str(row["vector"]) - if vector_str_raw.startswith("[") and vector_str_raw.endswith("]"): - vector_data = [float(x) for x in vector_str_raw[1:-1].split(",")] - - metadata = row["metadata"] if row["metadata"] else {} - if isinstance(metadata, str): - metadata = json.loads(metadata) - - metadata["score"] = 1 - distance - metadata["_distance"] = distance - - node = VectorNode( - vector_id=row["id"], - content=row["content"] or "", - vector=vector_data, - metadata=metadata, - ) - results.append(node) - - return results - - async def delete(self, vector_ids: str | list[str], **kwargs): - """Remove specific vector records from the collection by their IDs.""" - if isinstance(vector_ids, str): - vector_ids = [vector_ids] - - if not vector_ids: - return - - pool = await self._get_pool() - async with pool.acquire() as conn: - placeholders = ", ".join([f"${i + 1}" for i in range(len(vector_ids))]) - await conn.execute( - f"DELETE FROM {self.collection_name} WHERE id IN ({placeholders})", - *vector_ids, - ) - - logger.info(f"Deleted {len(vector_ids)} documents from {self.collection_name}") - - async def delete_all(self, **kwargs): - """Remove all vectors from the collection.""" - pool = await self._get_pool() - async with pool.acquire() as conn: - result = await conn.execute(f"DELETE FROM {self.collection_name}") - - logger.info(f"Deleted all documents from {self.collection_name} result={result}") - - async def update(self, nodes: VectorNode | list[VectorNode], **kwargs): - """Update existing vector nodes with new content, embeddings, or metadata.""" - if isinstance(nodes, VectorNode): - nodes = [nodes] - - if not nodes: - return - - nodes_without_vectors = [node for node in nodes if node.vector is None and node.content] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_update = [vector_map.get(n.vector_id, n) if n.vector is None and n.content else n for n in nodes] - else: - nodes_to_update = nodes - - pool = await self._get_pool() - async with pool.acquire() as conn: - for node in nodes_to_update: - update_fields = [] - params = [] - idx = 1 - - if node.content: - update_fields.append(f"content = ${idx}") - params.append(node.content) - idx += 1 - - if node.vector: - vector_str = f"[{','.join(map(str, node.vector))}]" - update_fields.append(f"vector = ${idx}::vector") - params.append(vector_str) - idx += 1 - - if node.metadata: - update_fields.append(f"metadata = ${idx}::jsonb") - params.append(json.dumps(node.metadata)) - idx += 1 - - if update_fields: - params.append(node.vector_id) - await conn.execute( - f"UPDATE {self.collection_name} SET {', '.join(update_fields)} WHERE id = ${idx}", - *params, - ) - - logger.info(f"Updated {len(nodes_to_update)} documents in {self.collection_name}") - - async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode] | None: - """Retrieve vector nodes by their unique identifiers.""" - single_result = isinstance(vector_ids, str) - if single_result: - vector_ids = [vector_ids] - - if not vector_ids: - return [] if not single_result else None - - pool = await self._get_pool() - async with pool.acquire() as conn: - placeholders = ", ".join([f"${i + 1}" for i in range(len(vector_ids))]) - rows = await conn.fetch( - f"SELECT id, content, vector, metadata FROM {self.collection_name} WHERE id IN ({placeholders})", - *vector_ids, - ) - - results = [] - for row in rows: - vector_data = None - if row["vector"]: - vector_str_raw = str(row["vector"]) - if vector_str_raw.startswith("[") and vector_str_raw.endswith("]"): - vector_data = [float(x) for x in vector_str_raw[1:-1].split(",")] - - metadata = row["metadata"] if row["metadata"] else {} - if isinstance(metadata, str): - metadata = json.loads(metadata) - - results.append( - VectorNode( - vector_id=row["id"], - content=row["content"] or "", - vector=vector_data, - metadata=metadata, - ), - ) - - if single_result: - return results[0] if results else None - return results - - async def list( - self, - filters: dict | None = None, - limit: int | None = None, - sort_key: str | None = None, - reverse: bool = False, - ) -> list[VectorNode]: - """Return a list of vector nodes matching the provided filters and limit. - - Args: - filters: Dictionary of filter conditions to match vectors - limit: Maximum number of vectors to return - sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting - reverse: If True, sort in descending order; if False, sort in ascending order - """ - pool = await self._get_pool() - filter_clause, filter_params = self._build_filter_clause(filters) - - # Add ORDER BY clause if sort_key is provided - order_clause = "" - if sort_key: - order_direction = "DESC" if reverse else "ASC" - order_clause = f"ORDER BY metadata->>'{sort_key}' {order_direction}" - - limit_clause = "" - if limit: - limit_clause = f"LIMIT ${len(filter_params) + 1}" - filter_params.append(limit) - - async with pool.acquire() as conn: - sql = f""" - SELECT id, content, vector, metadata - FROM {self.collection_name} - {filter_clause} - {order_clause} - {limit_clause} - """ - rows = await conn.fetch(sql, *filter_params) - - results = [] - for row in rows: - vector_data = None - if row["vector"]: - vector_str_raw = str(row["vector"]) - if vector_str_raw.startswith("[") and vector_str_raw.endswith("]"): - vector_data = [float(x) for x in vector_str_raw[1:-1].split(",")] - - metadata = row["metadata"] if row["metadata"] else {} - if isinstance(metadata, str): - metadata = json.loads(metadata) - - results.append( - VectorNode( - vector_id=row["id"], - content=row["content"] or "", - vector=vector_data, - metadata=metadata, - ), - ) - - return results - - async def collection_info(self) -> dict[str, Any]: - """Fetch metadata including record count and disk usage for the collection.""" - pool = await self._get_pool() - - async with pool.acquire() as conn: - row = await conn.fetchrow( - f""" - SELECT - '{self.collection_name}' as name, - (SELECT COUNT(*) FROM {self.collection_name}) as row_count, - pg_size_pretty(pg_total_relation_size('{self.collection_name}')) as total_size - """, - ) - - return { - "name": row["name"], - "count": row["row_count"], - "size": row["total_size"], - } - - async def reset(self): - """Purge all data by dropping and recreating the collection table.""" - logger.warning(f"Resetting collection {self.collection_name}...") - await self.delete_collection(self.collection_name) - await self.create_collection(self.collection_name) - - async def reset_collection(self, collection_name: str): - """Reset collection with table name validation for SQL injection prevention.""" - self._validate_table_name(collection_name) - self.collection_name = collection_name - await self.create_collection(collection_name) - logger.info(f"Collection reset to {collection_name}") - - async def start(self) -> None: - """Initialize the PGVector store. - - Creates the connection pool and ensures the collection table exists. - """ - await self._get_pool() - await super().start() - logger.info(f"PGVector collection {self.collection_name} initialized") - - async def close(self): - """Terminate the database connection pool and release associated resources.""" - if self._pool is not None: - await self._pool.close() - self._pool = None - logger.info("PGVector connection pool closed") diff --git a/reme/core/vector_store/qdrant_vector_store.py b/reme/core/vector_store/qdrant_vector_store.py deleted file mode 100644 index 3e69d667..00000000 --- a/reme/core/vector_store/qdrant_vector_store.py +++ /dev/null @@ -1,550 +0,0 @@ -"""Qdrant vector store implementation for the ReMe project.""" - -from pathlib import Path -from typing import Any - -from loguru import logger - -from .base_vector_store import BaseVectorStore -from ..embedding import BaseEmbeddingModel -from ..schema import VectorNode - -_QDRANT_IMPORT_ERROR: Exception | None = None - -try: - from qdrant_client import AsyncQdrantClient - from qdrant_client.models import ( - Distance, - FieldCondition, - Filter, - MatchValue, - PointIdsList, - PointStruct, - Range, - VectorParams, - ) -except Exception as e: - _QDRANT_IMPORT_ERROR = e - AsyncQdrantClient = None - Distance = None - FieldCondition = None - Filter = None - MatchValue = None - PointIdsList = None - PointStruct = None - Range = None - VectorParams = None - - -class QdrantVectorStore(BaseVectorStore): - """Vector store implementation using Qdrant for dense vector search.""" - - def __init__( - self, - collection_name: str, - db_path: str | Path, - embedding_model: BaseEmbeddingModel, - host: str | None = None, - port: int = 6333, - url: str | None = None, - api_key: str | None = None, - https: bool | None = None, - grpc_port: int = 6334, - prefer_grpc: bool = False, - distance: str = "cosine", - on_disk: bool = False, - **kwargs: Any, - ): - """Initialize the Qdrant client and collection configuration. - - Args: - collection_name: Name of the collection. - db_path: Local storage path for on-disk/in-memory mode. - embedding_model: Model used for generating vector embeddings. - host: Server host address. - port: HTTP port for the server. - url: Full connection URL. - api_key: Authentication key for Qdrant Cloud. - https: Use secure connection if True. - grpc_port: gRPC interface port. - prefer_grpc: Use gRPC instead of HTTP if True. - distance: Metric for similarity (cosine, euclid, dot). - on_disk: Enable persistent storage for vectors. - **kwargs: Additional client configuration. - """ - if _QDRANT_IMPORT_ERROR is not None: - raise ImportError( - "Qdrant requires extra dependencies. Install with `pip install qdrant-client`", - ) from _QDRANT_IMPORT_ERROR - - super().__init__( - collection_name=collection_name, - db_path=db_path, - embedding_model=embedding_model, - **kwargs, - ) - - self.is_local = host is None and url is None and api_key is None - self.client: AsyncQdrantClient - - # Store connection parameters for deferred initialization in start() - self._host = host - self._port = port - self._url = url - self._api_key = api_key - self._https = https - self._grpc_port = grpc_port - self._prefer_grpc = prefer_grpc - - distance_map = { - "cosine": Distance.COSINE, - "euclid": Distance.EUCLID, - "dot": Distance.DOT, - } - self.distance = distance_map.get(distance.lower(), Distance.COSINE) - self.on_disk = on_disk - - async def list_collections(self) -> list[str]: - """Retrieve names of all existing collections in the Qdrant instance.""" - collections = await self.client.get_collections() - return [collection.name for collection in collections.collections] - - async def create_collection(self, collection_name: str, **kwargs: Any): - """Create a new collection with the specified vector configuration. - - Args: - collection_name: Name of the collection to create. - **kwargs: Overrides for dimensions, distance, or on_disk settings. - """ - collections = await self.list_collections() - if collection_name in collections: - logger.info(f"Collection {collection_name} already exists") - return - - dimensions = kwargs.get("dimensions", self.embedding_model.dimensions) - distance = kwargs.get("distance", self.distance) - on_disk = kwargs.get("on_disk", self.on_disk) - - await self.client.create_collection( - collection_name=collection_name, - vectors_config=VectorParams( - size=dimensions, - distance=distance, - on_disk=on_disk, - ), - ) - - logger.info(f"Created collection {collection_name} with dimensions={dimensions}") - - if not self.is_local: - await self._create_payload_indexes(collection_name) - - async def _create_payload_indexes(self, collection_name: str): - """Create keyword indexes for common metadata fields to optimize filtering.""" - common_fields = ["user_id", "agent_id", "run_id", "actor_id", "source"] - - for field in common_fields: - try: - await self.client.create_payload_index( - collection_name=collection_name, - field_name=field, - field_schema="keyword", - ) - logger.debug(f"Created index for {field} in collection {collection_name}") - except Exception as e: - logger.debug(f"Index for {field} might already exist: {e}") - - async def delete_collection(self, collection_name: str, **kwargs: Any): - """Permanently remove a collection from the Qdrant instance.""" - collections = await self.list_collections() - if collection_name in collections: - await self.client.delete_collection(collection_name=collection_name) - logger.info(f"Deleted collection {collection_name}") - else: - logger.warning(f"Collection {collection_name} does not exist") - - async def copy_collection(self, collection_name: str, **kwargs: Any): - """Duplicate an existing collection to a new one including all data.""" - collection_info = await self.client.get_collection(collection_name=self.collection_name) - - await self.client.create_collection( - collection_name=collection_name, - vectors_config=collection_info.config.params.vectors, - ) - - offset = None - batch_size = 100 - - while True: - records, next_offset = await self.client.scroll( - collection_name=self.collection_name, - limit=batch_size, - offset=offset, - with_payload=True, - with_vectors=True, - ) - - if not records: - break - - points = [ - PointStruct( - id=record.id, - vector=record.vector, - payload=record.payload, - ) - for record in records - ] - - await self.client.upsert( - collection_name=collection_name, - points=points, - ) - - offset = next_offset - if offset is None: - break - - logger.info(f"Copied collection {self.collection_name} to {collection_name}") - - async def insert(self, nodes: VectorNode | list[VectorNode], **kwargs: Any): - """Insert vector nodes into the collection, generating embeddings as needed.""" - if isinstance(nodes, VectorNode): - nodes = [nodes] - - nodes_without_vectors = [node for node in nodes if node.vector is None] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_insert = [vector_map.get(n.vector_id, n) if n.vector is None else n for n in nodes] - else: - nodes_to_insert = nodes - - points = [] - for node in nodes_to_insert: - try: - point_id = int(node.vector_id) - except ValueError: - point_id = abs(hash(node.vector_id)) % (10**18) - - point = PointStruct( - id=point_id, - vector=node.vector, - payload={ - "vector_id": node.vector_id, - "content": node.content, - "metadata": node.metadata, - }, - ) - points.append(point) - - wait = kwargs.get("wait", True) - await self.client.upsert( - collection_name=self.collection_name, - points=points, - wait=wait, - ) - - logger.info(f"Inserted {len(points)} documents into {self.collection_name}") - - @staticmethod - def _create_filter(filters: dict) -> Filter | None: - """Convert a dictionary of filter conditions into a Qdrant Filter object. - - Supports two filter formats: - 1. Range query: {"field": [start_value, end_value]} - filters for field >= start_value AND field <= end_value - 2. Exact match: {"field": value} - filters for field == value - """ - if not filters: - return None - - conditions = [] - for key, value in filters.items(): - # New syntax: [start, end] represents a range query - if isinstance(value, list) and len(value) == 2: - # Range query: field >= value[0] AND field <= value[1] - # Qdrant's Range only supports numeric values - if isinstance(value[0], (int, float)) and isinstance(value[1], (int, float)): - conditions.append( - FieldCondition( - key=f"metadata.{key}", - range=Range(gte=value[0], lte=value[1]), - ), - ) - else: - # For non-numeric values (e.g., string dates), Qdrant doesn't support range queries - # We need to skip this filter with a warning - logger.warning( - f"Qdrant does not support range queries for non-numeric values. " - f"Skipping range filter for key '{key}' with values {value}. " - f"Consider using numeric timestamps instead.", - ) - elif isinstance(value, dict) and ("gte" in value or "lte" in value): - range_params = {} - # Check if values are numeric - if "gte" in value: - if isinstance(value["gte"], (int, float)): - range_params["gte"] = value["gte"] - else: - logger.warning( - f"Qdrant range filter for key '{key}' requires numeric gte value, " - f"got {type(value['gte']).__name__}. Skipping.", - ) - continue - if "lte" in value: - if isinstance(value["lte"], (int, float)): - range_params["lte"] = value["lte"] - else: - logger.warning( - f"Qdrant range filter for key '{key}' requires numeric lte value, " - f"got {type(value['lte']).__name__}. Skipping.", - ) - continue - - if range_params: # Only add condition if we have valid numeric parameters - conditions.append( - FieldCondition( - key=f"metadata.{key}", - range=Range(**range_params), - ), - ) - else: - # Exact match - conditions.append( - FieldCondition(key=f"metadata.{key}", match=MatchValue(value=value)), - ) - - return Filter(must=conditions) if conditions else None - - async def search( - self, - query: str, - limit: int = 5, - filters: dict | None = None, - **kwargs: Any, - ) -> list[VectorNode]: - """Search for the most similar vectors based on a text query.""" - query_vector = await self.get_embedding(query) - query_filter = self._create_filter(filters) if filters else None - score_threshold = kwargs.get("score_threshold", None) - - results = await self.client.query_points( - collection_name=self.collection_name, - query=query_vector, - query_filter=query_filter, - limit=limit, - score_threshold=score_threshold, - ) - - nodes = [] - for point in results.points: - payload = point.payload or {} - node = VectorNode( - vector_id=payload.get("vector_id", str(point.id)), - content=payload.get("content", ""), - vector=point.vector if hasattr(point, "vector") else None, - metadata=payload.get("metadata", {}), - ) - node.metadata["score"] = point.score - nodes.append(node) - - return nodes - - async def delete(self, vector_ids: str | list[str], **kwargs: Any): - """Delete specific vectors from the collection using their IDs.""" - if isinstance(vector_ids, str): - vector_ids = [vector_ids] - - point_ids = [] - for vector_id in vector_ids: - try: - point_id = int(vector_id) - except ValueError: - point_id = abs(hash(vector_id)) % (10**18) - point_ids.append(point_id) - - wait = kwargs.get("wait", True) - await self.client.delete( - collection_name=self.collection_name, - points_selector=PointIdsList(points=point_ids), - wait=wait, - ) - - logger.info(f"Deleted {len(point_ids)} documents from {self.collection_name}") - - async def delete_all(self, **kwargs: Any): - """Remove all vectors from the collection.""" - wait = kwargs.get("wait", True) - - # Delete all points by using an empty filter (matches all) - from qdrant_client.models import FilterSelector - - await self.client.delete( - collection_name=self.collection_name, - points_selector=FilterSelector(filter=Filter(must=[])), - wait=wait, - ) - - logger.info(f"Deleted all documents from {self.collection_name}") - - async def update(self, nodes: VectorNode | list[VectorNode], **kwargs: Any): - """Update existing vector nodes with new content or metadata.""" - if isinstance(nodes, VectorNode): - nodes = [nodes] - - nodes_without_vectors = [node for node in nodes if node.vector is None and node.content] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_update = [vector_map.get(n.vector_id, n) if n.vector is None and n.content else n for n in nodes] - else: - nodes_to_update = nodes - - points = [] - for node in nodes_to_update: - try: - point_id = int(node.vector_id) - except ValueError: - point_id = abs(hash(node.vector_id)) % (10**18) - - point = PointStruct( - id=point_id, - vector=node.vector, - payload={ - "vector_id": node.vector_id, - "content": node.content, - "metadata": node.metadata, - }, - ) - points.append(point) - - wait = kwargs.get("wait", True) - await self.client.upsert( - collection_name=self.collection_name, - points=points, - wait=wait, - ) - - logger.info(f"Updated {len(points)} documents in {self.collection_name}") - - async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode]: - """Retrieve vector nodes by their IDs from the collection.""" - single_result = isinstance(vector_ids, str) - if single_result: - vector_ids = [vector_ids] - - point_ids = [] - for vector_id in vector_ids: - try: - point_id = int(vector_id) - except ValueError: - point_id = abs(hash(vector_id)) % (10**18) - point_ids.append(point_id) - - points = await self.client.retrieve( - collection_name=self.collection_name, - ids=point_ids, - with_payload=True, - with_vectors=True, - ) - - results = [] - for point in points: - if point: - payload = point.payload or {} - node = VectorNode( - vector_id=payload.get("vector_id", str(point.id)), - content=payload.get("content", ""), - vector=point.vector, - metadata=payload.get("metadata", {}), - ) - results.append(node) - else: - logger.warning("Point not found") - - return results[0] if single_result and results else results - - async def list( - self, - filters: dict | None = None, - limit: int | None = None, - sort_key: str | None = None, - reverse: bool = False, - ) -> list[VectorNode]: - """List all vector nodes in the collection matching the filter criteria. - - Args: - filters: Dictionary of filter conditions to match vectors - limit: Maximum number of vectors to return - sort_key: Key to sort the results by (e.g., field name in metadata). None for no sorting - reverse: If True, sort in descending order; if False, sort in ascending order - """ - scroll_filter = self._create_filter(filters) if filters else None - - # If sorting is needed, fetch more records than the limit to ensure correct sorting - fetch_limit = 10000 if sort_key else (limit or 10000) - - records, _ = await self.client.scroll( - collection_name=self.collection_name, - scroll_filter=scroll_filter, - limit=fetch_limit, - with_payload=True, - with_vectors=True, - ) - - results = [] - for record in records: - payload = record.payload or {} - node = VectorNode( - vector_id=payload.get("vector_id", str(record.id)), - content=payload.get("content", ""), - vector=record.vector, - metadata=payload.get("metadata", {}), - ) - results.append(node) - - # Apply sorting if sort_key is provided - if sort_key: - # Sort with proper handling of None and missing values - def sort_key_func(node): - value = node.metadata.get(sort_key) - if value is None: - # Return appropriate default based on reverse flag - return float("-inf") if not reverse else float("inf") - return value - - results.sort(key=sort_key_func, reverse=reverse) - - # Apply limit after sorting - if limit is not None: - results = results[:limit] - - return results - - async def start(self) -> None: - """Initialize the Qdrant collection. - - Creates the collection if it doesn't exist with configured vector parameters. - For local mode, creates the db_path directory if it doesn't exist. - """ - if self.is_local: - self.db_path.mkdir(parents=True, exist_ok=True) - self.client = AsyncQdrantClient(path=str(self.db_path)) - else: - self.client = AsyncQdrantClient( - host=self._host, - port=self._port, - url=self._url, - api_key=self._api_key, - https=self._https, - grpc_port=self._grpc_port, - prefer_grpc=self._prefer_grpc, - **self.kwargs, - ) - await super().start() - logger.info(f"Qdrant collection {self.collection_name} initialized") - - async def close(self): - """Close the AsyncQdrantClient connection and release resources.""" - await self.client.close() - logger.info("Qdrant client connection closed") diff --git a/reme/core/vector_store/seekdb_vector_store.py b/reme/core/vector_store/seekdb_vector_store.py deleted file mode 100644 index 3e0bf24f..00000000 --- a/reme/core/vector_store/seekdb_vector_store.py +++ /dev/null @@ -1,437 +0,0 @@ -"""seekdb vector store implementation for the ReMe framework. - -Uses ``pyseekdb`` (Chroma-like Collection API) for **embedded** local storage or -**remote** OceanBase / seekdb—the same deployment modes as ``pyseekdb.Client``. -For SQL-table-oriented helpers via ``pyobvector``, see ``ObVecVectorStore``. -""" - -from pathlib import Path -from typing import Any - -from loguru import logger - -from .base_vector_store import BaseVectorStore -from ..embedding import BaseEmbeddingModel -from ..schema import VectorNode -from ..utils.pyseekdb_conn import ( - DEFAULT_SEEKDB_DATABASE, - admin_kwargs_from_client_kwargs, - build_pyseekdb_client_kwargs, -) - -# Optional: preserve original exception for "raise ... from _PYSEEKDB_IMPORT_ERROR" (better diagnostics) -_PYSEEKDB_IMPORT_ERROR = None - -try: - import pyseekdb - from pyseekdb import Configuration, HNSWConfiguration - - PYSEEKDB_AVAILABLE = True -except ImportError as e: - _PYSEEKDB_IMPORT_ERROR = e - pyseekdb = None - Configuration = None - HNSWConfiguration = None - - -class SeekdbVectorStore(BaseVectorStore): - """Vector store using ``pyseekdb`` and the Chroma-like Collection API. - - **Embedded** (default): optional ``path`` to the embedded data directory; if omitted, - pyseekdb applies its default (typically a ``seekdb.db`` directory name). **Remote**: - ``host`` / ``port`` plus auth (same deployment style as ``ObVecVectorStore``, without ``uri``). - - Vector similarity search and metadata filtering; no full-text index by default. - """ - - def __init__( - self, - collection_name: str, - db_path: str | Path, - embedding_model: BaseEmbeddingModel, - database: str = DEFAULT_SEEKDB_DATABASE, - distance: str = "cosine", - host: str | None = None, - port: int | None = None, - user: str | None = None, - password: str = "", - path: str | None = None, - **kwargs: Any, - ): - """Initialize the seekdb vector store. - - Args: - collection_name: Name of the collection. - db_path: Working directory for ReMe (metadata sidecar); also used when resolving - a default location alongside **remote** mode (mirrors ``ObVecVectorStore``). - embedding_model: Model used for generating vector embeddings. - database: Database name on the seekdb / OceanBase instance. - distance: Similarity metric: cosine, euclid, dot. - host: Remote server host (embedded mode if unset or empty). - port: Remote port (default ``2881`` when ``host`` is set). - user: Remote user (``None`` uses library default ``root``). - password: Remote password. - path: Embedded data directory passed to ``pyseekdb.Client``; omit to use the - library default (typically ``./seekdb.db`` as the directory name). - **kwargs: Additional options (ignored for compatibility). - """ - if _PYSEEKDB_IMPORT_ERROR is not None: - raise ImportError( - "seekdb vector store requires pyseekdb. Install with `pip install pyseekdb` or `pip install reme-ai`", - ) from _PYSEEKDB_IMPORT_ERROR - - super().__init__( - collection_name=collection_name, - db_path=db_path, - embedding_model=embedding_model, - **kwargs, - ) - self.database = database - self.distance = distance.lower() - self.client: Any = None - self.collection: Any = None - - self._is_remote, self._client_kw = build_pyseekdb_client_kwargs( - path=None if (host and host.strip()) else path, - database=self.database, - host=host, - port=port, - user=user, - password=password, - ) - - def _client_kwargs(self) -> dict: - """Kwargs passed to ``pyseekdb.Client`` (embedded or remote).""" - return self._client_kw - - def _coerce_embedding_for_upsert(self, vec: Any) -> list[float]: - """Normalize vectors before ``collection.upsert`` (pyseekdb SQL rejects empty hex).""" - if vec is None: - raw: list[float] = [] - elif hasattr(vec, "tolist"): - raw = list(vec.tolist()) - elif isinstance(vec, list): - raw = vec - else: - raw = list(vec) - dim = self.embedding_model.dimensions - actual_len = len(raw) - if actual_len == dim: - return raw - if actual_len < dim: - logger.warning( - f"Embedding dimensions {actual_len} < {dim}, padding with zeros", - ) - return raw + [0.0] * (dim - actual_len) - logger.warning(f"Embedding dimensions {actual_len} > {dim}, truncating") - return raw[:dim] - - @staticmethod - def _build_where(filters: dict | None) -> dict | None: - """Build seekdb/Chroma-style where clause from universal filter format. - - Supports exact match and range: {"key": value} or {"key": [start, end]}. - """ - if not filters: - return None - conditions = [] - for key, value in filters.items(): - if value == "*": - continue - if isinstance(value, list) and len(value) == 2: - conditions.append({key: {"$gte": value[0]}}) - conditions.append({key: {"$lte": value[1]}}) - elif isinstance(value, dict) and ("gte" in value or "lte" in value or "gt" in value or "lt" in value): - for op, val in value.items(): - if op in ("gte", "lte", "gt", "lt") and val is not None: - conditions.append({key: {"$" + op: val}}) - else: - conditions.append({key: {"$eq": value}}) - if not conditions: - return None - return conditions[0] if len(conditions) == 1 else {"$and": conditions} - - @staticmethod - def _parse_results( - ids: list, - documents: list | None = None, - metadatas: list | None = None, - embeddings: list | None = None, - distances: list | None = None, - include_score: bool = False, - ) -> list[VectorNode]: - """Convert seekdb get/query result to list of VectorNode.""" - nodes = [] - documents = documents or [] - metadatas = metadatas or [] - embeddings = embeddings or [] - distances = distances or [] - if ids and isinstance(ids, list) and ids and isinstance(ids[0], list): - ids = ids[0] - if documents and isinstance(documents[0], list): - documents = documents[0] - if metadatas and isinstance(metadatas[0], list): - metadatas = metadatas[0] - if embeddings and isinstance(embeddings[0], list): - embeddings = embeddings[0] - if distances and isinstance(distances[0], (list, tuple)): - distances = distances[0] - for i, vector_id in enumerate(ids): - meta = metadatas[i] if i < len(metadatas) and metadatas[i] is not None else {} - if include_score and i < len(distances): - meta = dict(meta) - meta["score"] = 1.0 - (float(distances[i]) / 2.0) if distances[i] is not None else 0.0 - nodes.append( - VectorNode( - vector_id=str(vector_id), - content=documents[i] if i < len(documents) and documents[i] is not None else "", - vector=embeddings[i] if i < len(embeddings) else None, - metadata=meta, - ), - ) - return nodes - - async def list_collections(self) -> list[str]: - """List collection names in the current database.""" - if self.client is None: - return [] - try: - colls = self.client.list_collections() - return [c.name if hasattr(c, "name") else str(c) for c in colls] - except Exception as e: - logger.debug("seekdb list_collections: %s", e) - return [self.collection_name] - - async def create_collection(self, collection_name: str, **kwargs: Any) -> None: - """Create or get collection with HNSW vector index.""" - if self.client is None: - raise RuntimeError("seekdb client not initialized; call start() first") - dimensions = kwargs.get("dimensions", self.embedding_model.dimensions) - config = Configuration( - hnsw=HNSWConfiguration(dimension=dimensions, distance=self.distance), - ) - coll = self.client.get_or_create_collection( - name=collection_name, - configuration=config, - embedding_function=None, - ) - if collection_name == self.collection_name: - self.collection = coll - logger.info(f"seekdb collection {collection_name} ready (dim={dimensions})") - - async def delete_collection(self, collection_name: str, **kwargs: Any) -> None: - """Remove the collection from the database.""" - if self.client is None: - return - try: - self.client.delete_collection(collection_name) - if collection_name == self.collection_name: - self.collection = None - logger.info(f"Deleted seekdb collection {collection_name}") - except Exception as e: - logger.warning("seekdb delete_collection %s: %s", collection_name, e) - - async def copy_collection(self, collection_name: str, **kwargs: Any) -> None: - """Copy current collection to a new one.""" - if self.collection is None: - raise RuntimeError("No current collection") - data = self.collection.get(include=["documents", "metadatas", "embeddings"]) - ids = data.get("ids") or [] - if not ids: - logger.warning("Source collection is empty") - return - dims = self.embedding_model.dimensions - embs = data.get("embeddings") - if embs and (isinstance(embs[0], list) and embs[0]) or (not isinstance(embs[0], list) and embs): - dims = len(embs[0]) if isinstance(embs[0], list) else len(embs) - config = Configuration( - hnsw=HNSWConfiguration(dimension=dims, distance=self.distance), - ) - self.client.get_or_create_collection( - name=collection_name, - configuration=config, - embedding_function=None, - ) - new_coll = self.client.get_collection(name=collection_name, embedding_function=None) - emb_out = data.get("embeddings") or [] - emb_norm = [self._coerce_embedding_for_upsert(e) for e in emb_out] if emb_out else [] - new_coll.upsert( - ids=ids, - documents=data.get("documents", []), - embeddings=emb_norm, - metadatas=data.get("metadatas", []), - ) - logger.info(f"Copied {self.collection_name} to {collection_name}") - - async def insert(self, nodes: VectorNode | list[VectorNode], **kwargs: Any) -> None: - """Insert vector nodes; generate embeddings for nodes that lack them.""" - if isinstance(nodes, VectorNode): - nodes = [nodes] - if not nodes: - return - nodes_without_vectors = [n for n in nodes if n.vector is None] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_insert = [vector_map.get(n.vector_id, n) if n.vector is None else n for n in nodes] - else: - nodes_to_insert = nodes - ids = [n.vector_id for n in nodes_to_insert] - documents = [n.content for n in nodes_to_insert] - embeddings = [self._coerce_embedding_for_upsert(n.vector) for n in nodes_to_insert] - metadatas = [n.metadata for n in nodes_to_insert] - self.collection.upsert(ids=ids, documents=documents, embeddings=embeddings, metadatas=metadatas) - logger.info(f"Inserted {len(nodes_to_insert)} nodes into {self.collection_name}") - - async def search( - self, - query: str, - limit: int = 5, - filters: dict | None = None, - **kwargs: Any, - ) -> list[VectorNode]: - """Vector similarity search with optional metadata filter.""" - query_vector = await self.get_embedding(query) - where = self._build_where(filters) - results = self.collection.query( - query_embeddings=[query_vector], - n_results=limit, - where=where, - include=["documents", "metadatas", "distances"], - ) - ids = results.get("ids") or [] - documents = results.get("documents") - metadatas = results.get("metadatas") - distances = results.get("distances") - nodes = self._parse_results( - ids, - documents=documents, - metadatas=metadatas, - embeddings=results.get("embeddings"), - distances=distances, - include_score=True, - ) - score_threshold = kwargs.get("score_threshold") - if score_threshold is not None: - nodes = [n for n in nodes if n.metadata.get("score", 0) >= score_threshold] - return nodes - - async def delete(self, vector_ids: str | list[str], **kwargs: Any) -> None: - """Delete points by id.""" - if isinstance(vector_ids, str): - vector_ids = [vector_ids] - if not vector_ids: - return - self.collection.delete(ids=vector_ids) - logger.info(f"Deleted {len(vector_ids)} nodes from {self.collection_name}") - - async def delete_all(self, **kwargs: Any) -> None: - """Remove all points from the collection.""" - data = self.collection.get(include=[]) - ids = data.get("ids") or [] - if ids: - self.collection.delete(ids=ids) - logger.info(f"Deleted all {len(ids)} nodes from {self.collection_name}") - - async def update(self, nodes: VectorNode | list[VectorNode], **kwargs: Any) -> None: - """Update nodes (upsert by id).""" - if isinstance(nodes, VectorNode): - nodes = [nodes] - if not nodes: - return - nodes_without_vectors = [n for n in nodes if n.vector is None and n.content] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_update = [vector_map.get(n.vector_id, n) if n.vector is None and n.content else n for n in nodes] - else: - nodes_to_update = nodes - ids = [n.vector_id for n in nodes_to_update] - documents = [n.content for n in nodes_to_update] - embeddings = [self._coerce_embedding_for_upsert(n.vector) for n in nodes_to_update] - metadatas = [n.metadata for n in nodes_to_update] - self.collection.upsert(ids=ids, documents=documents, embeddings=embeddings, metadatas=metadatas) - logger.info(f"Updated {len(nodes_to_update)} nodes in {self.collection_name}") - - async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode]: - """Fetch nodes by id.""" - single = isinstance(vector_ids, str) - ids = [vector_ids] if single else list(vector_ids) - if not ids: - return None if single else [] - results = self.collection.get( - ids=ids, - include=["documents", "metadatas", "embeddings"], - ) - rids = results.get("ids") or [] - nodes = self._parse_results( - rids, - documents=results.get("documents"), - metadatas=results.get("metadatas"), - embeddings=results.get("embeddings"), - ) - if single: - return nodes[0] if nodes else None - return nodes - - async def list( - self, - filters: dict | None = None, - limit: int | None = None, - sort_key: str | None = None, - reverse: bool = True, - ) -> list[VectorNode]: - """List nodes with optional filter, limit, and sort by metadata key.""" - where = self._build_where(filters) - # When sorting in memory, fetch candidates first (cap like default list); do not pass - # user limit to get() or we sort an arbitrary first page only (see ChromaVectorStore.list). - if sort_key is not None: - fetch_limit = 10000 - else: - fetch_limit = limit if limit is not None else 10000 - results = self.collection.get( - where=where, - limit=fetch_limit, - include=["documents", "metadatas", "embeddings"], - ) - ids = results.get("ids") or [] - nodes = self._parse_results( - ids, - documents=results.get("documents"), - metadatas=results.get("metadatas"), - embeddings=results.get("embeddings"), - ) - if sort_key: - - def key_fn(n): - v = n.metadata.get(sort_key) - if v is None: - return float("-inf") if not reverse else float("inf") - return v - - nodes.sort(key=key_fn, reverse=reverse) - if limit is not None: - nodes = nodes[:limit] - return nodes - - async def start(self) -> None: - """Initialize seekdb client and ensure collection exists.""" - kw = self._client_kwargs() - if not self._is_remote and "path" in kw: - Path(kw["path"]).parent.mkdir(parents=True, exist_ok=True) - try: - admin = pyseekdb.AdminClient(**admin_kwargs_from_client_kwargs(kw)) - if not any(db.name == self.database for db in admin.list_databases()): - admin.create_database(self.database) - except Exception as e: - logger.debug("seekdb AdminClient create_database: %s", e) - self.client = pyseekdb.Client(**kw) - await self.create_collection(self.collection_name) - mode = "remote" if self._is_remote else "embedded" - logger.info(f"seekdb vector store ({mode}) {self.collection_name} initialized") - - async def close(self) -> None: - """Release client; no explicit close in pyseekdb, clear references.""" - self.client = None - self.collection = None - logger.info("seekdb vector store closed") diff --git a/reme/core/vector_store/zvec_vector_store.py b/reme/core/vector_store/zvec_vector_store.py deleted file mode 100644 index dde12a0b..00000000 --- a/reme/core/vector_store/zvec_vector_store.py +++ /dev/null @@ -1,809 +0,0 @@ -"""Zvec vector store implementation for the ReMe framework.""" - -from __future__ import annotations - -import json -from pathlib import Path -from typing import Any - -from loguru import logger - -from .base_vector_store import BaseVectorStore -from ..embedding import BaseEmbeddingModel -from ..schema import VectorNode - -_ZVEC_IMPORT_ERROR: Exception | None = None - -try: - import zvec # type: ignore[import-untyped] - from zvec import ( - CollectionOption, - CollectionSchema, - DataType, - Doc, - FieldSchema, - HnswIndexParam, - InvertIndexParam, - VectorQuery, - VectorSchema, - ) - from zvec.typing import MetricType -except Exception as e: - _ZVEC_IMPORT_ERROR = e - zvec = None # type: ignore[assignment] - - -# Default vector field name used inside zvec collections -_DEFAULT_VECTOR_FIELD = "embedding" - -# Default scalar content field name for storing text -_CONTENT_FIELD = "content" - -# Field name for JSON-serialized metadata -_METADATA_FIELD = "metadata" - -# Metadata fields promoted to top-level zvec schema columns for native filtering. -# These are the most commonly filtered keys in ReMe's memory system. -# Defining them as independent schema columns allows zvec to perform -# filtering at the database level instead of Python post-filtering. -# Format: {metadata_key: (zvec_data_type_str, has_inverted_index)} -_PROMOTED_FIELD_SPECS: dict[str, tuple[str, bool]] = { - "memory_type": ("STRING", True), # Inverted index for exact match filtering - "memory_target": ("STRING", True), # Inverted index for exact match filtering - "author": ("STRING", False), - "time_int": ("INT64", False), # Numeric for range queries -} - -# zvec data-type string → DataType enum mapping (populated after import) -_DATATYPE_MAP: dict[str, Any] = {} # filled in _build_collection_schema - - -def _escape_zvec_string(value: str) -> str: - """Escape a string value for use in zvec filter expressions.""" - return value.replace("'", "\\'") - - -def _build_zvec_filter( - filters: dict | None, - promoted_fields: set[str], -) -> tuple[str | None, dict | None]: - """Split ReMe filter dict into a zvec native filter expression and remaining post-filters. - - For filter keys that correspond to promoted schema fields, native - zvec filter expressions are generated. Non-promoted keys are - kept for Python post-filtering. - - Args: - filters: ReMe-style filter dictionary. - promoted_fields: Set of metadata keys that exist as top-level schema columns. - - Returns: - (native_filter_expr, post_filter_dict) — either may be None. - """ - if not filters: - return None, None - - native_conditions: list[str] = [] - post_filters: dict = {} - - for key, value in filters.items(): - if key.startswith("$"): - # Compound operators ($or, $and, $not) — keep for post-filtering - post_filters[key] = value - continue - - if key not in promoted_fields: - # Not a promoted field — use post-filtering - post_filters[key] = value - continue - - # Build native filter condition for promoted fields - field_type = _PROMOTED_FIELD_SPECS.get(key, ("STRING", False))[0] - - if isinstance(value, list) and len(value) == 2: - # Range query: [start, end] - if field_type == "INT64": - native_conditions.append(f"{key} >= {value[0]} AND {key} <= {value[1]}") - else: - # STRING range — use >= and <= with string escaping - native_conditions.append( - f"{key} >= '{_escape_zvec_string(str(value[0]))}' " - f"AND {key} <= '{_escape_zvec_string(str(value[1]))}'", - ) - elif isinstance(value, bool): - native_conditions.append(f"{key} = {str(value).upper()}") - elif isinstance(value, (int, float)): - native_conditions.append(f"{key} = {value}") - elif isinstance(value, str): - native_conditions.append(f"{key} = '{_escape_zvec_string(value)}'") - else: - # Unsupported type — fall back to post-filtering - post_filters[key] = value - - native_filter = " AND ".join(native_conditions) if native_conditions else None - return native_filter, post_filters if post_filters else None - - -def _metric_type_from_str(metric: str) -> Any: - """Convert a string metric name to zvec MetricType enum value.""" - if zvec is None: - return None - mapping = { - "cosine": MetricType.COSINE, - "l2": MetricType.L2, - "ip": MetricType.IP, - } - return mapping.get(metric.lower(), MetricType.COSINE) - - -def _build_collection_schema( - name: str, - dimension: int, - metric: str = "cosine", -) -> CollectionSchema: - """Build a zvec CollectionSchema for ReMe usage. - - The schema contains: - - "content" (STRING, inverted index) — text content - - "metadata" (STRING) — JSON-serialized metadata dictionary - - Promoted metadata fields (STRING / INT64) — for native zvec filtering - - "embedding" (VECTOR_FP32, dimension, HNSW index) — the vector field - - Promoted fields are commonly filtered metadata keys defined as top-level - schema columns so that zvec can perform filtering natively instead of - Python post-filtering. The full metadata is still stored as JSON in the - "metadata" field for complete round-trip serialization. - - zvec automatically manages the document ID (string type); we do NOT - define an "id" field in the schema. - """ - # Populate the DataType map on first call - if not _DATATYPE_MAP: - _DATATYPE_MAP.update( - { - "STRING": DataType.STRING, - "INT64": DataType.INT64, - }, - ) - - distance = _metric_type_from_str(metric) - - # Base fields - fields = [ - FieldSchema("content", DataType.STRING, nullable=True, index_param=InvertIndexParam()), - FieldSchema("metadata", DataType.STRING, nullable=True), - ] - - # Add promoted metadata fields as top-level schema columns - for field_name, (type_str, has_inv_index) in _PROMOTED_FIELD_SPECS.items(): - dt = _DATATYPE_MAP[type_str] - idx_param = InvertIndexParam() if has_inv_index else None - fields.append(FieldSchema(field_name, dt, nullable=True, index_param=idx_param)) - - return CollectionSchema( - name=name, - fields=fields, - vectors=[ - VectorSchema( - name=_DEFAULT_VECTOR_FIELD, - data_type=DataType.VECTOR_FP32, - dimension=dimension, - index_param=HnswIndexParam(metric_type=distance), - ), - ], - ) - - -def _vector_node_to_doc(node: VectorNode) -> Doc: - """Convert a ReMe VectorNode to a zvec Doc. - - Metadata is serialized as a JSON string into the "metadata" field. - The "score" key is excluded since it is a computed value, not stored data. - Promoted metadata fields are also extracted as top-level Doc fields - for native zvec filtering. - The vector is placed under the default vector field name. - The zvec Doc id must be a string. - """ - # Filter out computed score before serialization - meta_to_store = {k: v for k, v in node.metadata.items() if k != "score"} - - fields: dict[str, Any] = { - "content": node.content, - "metadata": json.dumps(meta_to_store) if meta_to_store else "{}", - } - - # Extract promoted metadata fields as top-level schema columns - for field_name, (type_str, _) in _PROMOTED_FIELD_SPECS.items(): - value = meta_to_store.get(field_name) - if value is not None: - # Ensure correct type: INT64 fields must be int - if type_str == "INT64" and not isinstance(value, int): - try: - value = int(value) - except (ValueError, TypeError): - continue - fields[field_name] = value - - vectors: dict[str, Any] = {} - if node.vector is not None: - vectors[_DEFAULT_VECTOR_FIELD] = node.vector - - return Doc(id=str(node.vector_id), fields=fields, vectors=vectors) - - -def _doc_to_vector_node(doc: Doc, include_score: bool = False) -> VectorNode: - """Convert a zvec Doc back to a ReMe VectorNode. - - The "metadata" field is parsed from JSON. The "content" field becomes - the node content. If ``include_score`` is True, the search score is - added to the metadata dictionary. - """ - metadata: dict[str, str | bool | int | float] = {} - - # Parse JSON metadata - raw_metadata = doc.field("metadata") - if raw_metadata: - try: - parsed = json.loads(raw_metadata) - if isinstance(parsed, dict): - metadata.update(parsed) - except (json.JSONDecodeError, TypeError): - logger.warning(f"Failed to parse metadata JSON: {raw_metadata}") - - if include_score and doc.score is not None: - metadata["score"] = doc.score - - # Extract vector — doc.vector() returns list or empty dict - raw_vector = doc.vector(_DEFAULT_VECTOR_FIELD) - vector = raw_vector if isinstance(raw_vector, list) and len(raw_vector) > 0 else None - - content = doc.field("content") or "" - - return VectorNode( - vector_id=str(doc.id), - content=str(content), - vector=vector, - metadata=metadata, - ) - - -def _apply_filters_post(nodes: list[VectorNode], filters: dict | None) -> list[VectorNode]: - """Apply ReMe-style filter dict as post-filtering on metadata. - - Used as a fallback for metadata keys that are NOT promoted to top-level - schema columns (and thus cannot be filtered natively by zvec). Promoted - fields are handled by zvec's native ``filter`` parameter instead. - - Supports: - - Exact match: {"field": value} - - Range query: {"field": [start, end]} - """ - if not filters: - return nodes - - filtered = [] - for node in nodes: - match = True - for key, value in filters.items(): - if key.startswith("$"): - # Skip compound operators for post-filtering - continue - node_value = node.metadata.get(key) - - # Range query: [start, end] - if isinstance(value, list) and len(value) == 2: - if node_value is None: - match = False - break - try: - if not value[0] <= node_value <= value[1]: - match = False - break - except TypeError: - match = False - break - else: - # Exact match - if node_value != value: - match = False - break - - if match: - filtered.append(node) - - return filtered - - -class ZvecVectorStore(BaseVectorStore): - """Zvec-based vector store implementation. - - Zvec is a high-performance vector database. This adapter bridges the - ReMe ``BaseVectorStore`` interface with zvec's Python API. - - Supports local persistent storage via ``db_path``. - - Args: - collection_name: Name of the vector collection. - db_path: Local storage path for persistent mode. - embedding_model: Model used for generating vector embeddings. - dimension: Dimensionality of the embedding vectors (default: 1024). - distance: Distance metric — cosine / l2 / ip (default: cosine). - **kwargs: Additional zvec-specific configuration. - """ - - def __init__( - self, - collection_name: str, - db_path: str | Path, - embedding_model: BaseEmbeddingModel, - dimension: int = 1024, - distance: str = "cosine", - **kwargs: Any, - ): - """Initialize the Zvec vector store.""" - if _ZVEC_IMPORT_ERROR is not None: - raise ImportError( - "Zvec requires extra dependencies. Install with `pip install zvec`", - ) from _ZVEC_IMPORT_ERROR - - super().__init__( - collection_name=collection_name, - db_path=db_path, - embedding_model=embedding_model, - **kwargs, - ) - - self.dimension = dimension - self.distance = distance - self._collection = None - self._initialized = False - # Set of promoted field names that exist in the current collection's schema. - # Populated during start() by inspecting the schema. Only fields present - # in the schema can use native zvec filtering; the rest fall back to - # Python post-filtering. - self._promoted_fields_in_schema: set[str] = set() - - # ------------------------------------------------------------------ - # Lifecycle - # ------------------------------------------------------------------ - - async def start(self) -> None: - """Initialize the Zvec engine and open the collection. - - Calls ``zvec.init()`` once, then tries to ``zvec.open()`` an existing - collection or ``zvec.create_and_open()`` a new one. - After opening, detects which promoted fields exist in the schema - and attempts to add missing numeric fields via ``add_column``. - """ - if not self._initialized: - try: - zvec.init() - except RuntimeError: - # Already initialized — safe to ignore - pass - self._initialized = True - - self.db_path.mkdir(parents=True, exist_ok=True) - collection_path = str(self.db_path / self.collection_name) - - option = CollectionOption(read_only=False, enable_mmap=True) - - try: - # Try opening an existing collection first - self._collection = zvec.open(collection_path, option) - logger.info(f"Opened existing Zvec collection at {collection_path}") - except Exception: - # Collection doesn't exist — create it - schema = _build_collection_schema( - name=self.collection_name, - dimension=self.dimension, - metric=self.distance, - ) - self._collection = zvec.create_and_open( - path=collection_path, - schema=schema, - option=option, - ) - logger.info(f"Created new Zvec collection at {collection_path}") - - # Detect which promoted fields exist in the current schema - self._detect_promoted_fields() - - # Try to add missing numeric promoted fields to existing collections - # (zvec's add_column only supports numeric types: INT64, FLOAT, etc.) - self._ensure_numeric_promoted_columns() - - async def close(self) -> None: - """Flush pending writes and release the collection handle.""" - if self._collection is not None: - try: - self._collection.flush() - except Exception as e: - logger.warning(f"Failed to flush collection on close: {e}") - self._collection = None - logger.info(f"Zvec vector store for collection {self.collection_name} closed") - - # ------------------------------------------------------------------ - # Collection management - # ------------------------------------------------------------------ - - async def list_collections(self) -> list[str]: - """Retrieve a list of collection names in the db_path directory. - - Zvec doesn't have a global ``list_collections`` API; we scan the - db_path directory for zvec collection folders. - """ - if not self.db_path.exists(): - return [] - collections = [] - for child in self.db_path.iterdir(): - if child.is_dir(): - collections.append(child.name) - return collections - - async def create_collection(self, collection_name: str, **kwargs) -> None: - """Create a new collection with the specified name and distance metric.""" - if not self._initialized: - try: - zvec.init() - except RuntimeError: - pass - self._initialized = True - - self.db_path.mkdir(parents=True, exist_ok=True) - collection_path = str(self.db_path / collection_name) - - dimension = kwargs.get("dimension", self.dimension) - metric = kwargs.get("distance_metric", self.distance) - - schema = _build_collection_schema( - name=collection_name, - dimension=dimension, - metric=metric, - ) - option = CollectionOption(read_only=False, enable_mmap=True) - - collection = zvec.create_and_open(path=collection_path, schema=schema, option=option) - if collection_name == self.collection_name: - self._collection = collection - logger.info(f"Created collection `{collection_name}`") - - async def delete_collection(self, collection_name: str, **kwargs) -> None: - """Permanently remove a collection from disk.""" - # If it's the active collection, destroy it via zvec API - if self._collection is not None and collection_name == self.collection_name: - try: - self._collection.destroy() - self._collection = None - deleted = True - except Exception as _e: - logger.warning(f"Failed to destroy collection {collection_name}: {_e}") - deleted = False - else: - # For non-active collections, remove the directory - collection_path = self.db_path / collection_name - if collection_path.exists(): - import shutil - - shutil.rmtree(collection_path, ignore_errors=True) - deleted = True - else: - deleted = False - - logger.info(f"Deleted collection {collection_name}: {deleted}") - - async def copy_collection(self, collection_name: str, **kwargs) -> None: - """Duplicate the current collection to a new one with the given name. - - Uses ``shutil.copytree`` to directly copy the collection directory on - disk, which is both faster and complete — it avoids the topk limit of - ``list()`` (max 1024 docs) that would cause data loss for large - collections. - - The source collection is flushed before copying to ensure all - pending writes are persisted to disk. - """ - import shutil - - # Flush source collection so all data is on disk - if self._collection is not None: - self._collection.flush() - - src_path = self.db_path / self.collection_name - dst_path = self.db_path / collection_name - - if not src_path.exists(): - logger.warning(f"Source collection directory not found: {src_path}") - return - - if dst_path.exists(): - logger.warning(f"Target collection already exists: {dst_path}, removing it first") - shutil.rmtree(dst_path, ignore_errors=True) - - shutil.copytree(src_path, dst_path) - logger.info( - f"Copied collection {self.collection_name} to {collection_name} " - f"(directory copy: {src_path} -> {dst_path})", - ) - - # ------------------------------------------------------------------ - # CRUD operations - # ------------------------------------------------------------------ - - async def insert(self, nodes: VectorNode | list[VectorNode], **kwargs) -> None: - """Add one or more vector nodes into the current collection. - - Automatically generates embeddings for nodes that lack vectors. - """ - if isinstance(nodes, VectorNode): - nodes = [nodes] - if not nodes: - return - - # Batch generate embeddings for nodes that need them - nodes_without_vectors = [n for n in nodes if n.vector is None] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_insert = [vector_map.get(n.vector_id, n) if n.vector is None else n for n in nodes] - else: - nodes_to_insert = nodes - - batch_size = kwargs.get("batch_size", 100) - - for i in range(0, len(nodes_to_insert), batch_size): - batch = nodes_to_insert[i : i + batch_size] - docs = [_vector_node_to_doc(n) for n in batch] - self._collection.insert(docs) - - logger.info(f"Inserted {len(nodes_to_insert)} nodes into {self.collection_name}") - - async def search( - self, - query: str, - limit: int = 5, - filters: dict | None = None, - **kwargs, - ) -> list[VectorNode]: - """Find the most similar vector nodes based on a text query. - - Uses zvec's ``query()`` method with a ``VectorQuery`` built from the - embedding of the query text. Promoted metadata fields are filtered - natively via zvec's ``filter`` parameter; remaining filters are - applied as post-filtering in Python. - """ - query_vector = await self.get_embedding(query) - - vq = VectorQuery( - field_name=_DEFAULT_VECTOR_FIELD, - vector=query_vector, - ) - include_vector = kwargs.get("include_embeddings", False) - - # Split filters: native zvec filter vs Python post-filter - native_filter, post_filters = _build_zvec_filter(filters, self._promoted_fields_in_schema) - - # Over-fetch to compensate for post-filtering - _ZVEC_MAX_TOPK = 1024 - # When post-filters remain, we need to fetch more results because - # many may be filtered out. Use the maximum allowed to minimize misses. - fetch_limit = _ZVEC_MAX_TOPK if post_filters else min(limit, _ZVEC_MAX_TOPK) - - results = self._collection.query( - vectors=vq, - topk=fetch_limit, - filter=native_filter, - include_vector=include_vector, - ) - - nodes = [_doc_to_vector_node(doc, include_score=True) for doc in results] - - # Post-filter on non-promoted metadata fields - nodes = _apply_filters_post(nodes, post_filters) - - score_threshold = kwargs.get("score_threshold") - if score_threshold is not None: - nodes = [n for n in nodes if n.metadata.get("score", 0) >= score_threshold] - - return nodes[:limit] - - async def delete(self, vector_ids: str | list[str], **kwargs) -> None: - """Remove specific vectors from the collection using their identifiers.""" - if isinstance(vector_ids, str): - vector_ids = [vector_ids] - if not vector_ids: - return - - self._collection.delete(vector_ids) - logger.info(f"Deleted {len(vector_ids)} nodes from {self.collection_name}") - - async def delete_all(self, **kwargs) -> None: - """Remove all vectors from the collection. - - Uses zvec's ``delete_by_filter`` with a condition that matches all - documents (content is not empty), or falls back to query + delete - in batches (zvec topk max is 1024). - """ - stats = self._collection.stats - count = stats.doc_count if stats else 0 - if count > 0: - try: - # Use delete_by_filter for efficiency - self._collection.delete_by_filter("content!=''") - except Exception: - # Fallback: fetch all IDs in batches then delete - _ZVEC_MAX_TOPK = 1024 - remaining = count - while remaining > 0: - all_docs = self._collection.query(topk=min(remaining, _ZVEC_MAX_TOPK), include_vector=False) - if not all_docs: - break - ids = [doc.id for doc in all_docs] - self._collection.delete(ids) - remaining -= len(ids) - logger.info(f"Deleted all {count} nodes from {self.collection_name}") - - async def update(self, nodes: VectorNode | list[VectorNode], **kwargs) -> None: - """Update existing vectors using zvec's ``upsert``. - - Automatically regenerates embeddings for nodes whose content changed - but lack an updated vector. - """ - if isinstance(nodes, VectorNode): - nodes = [nodes] - if not nodes: - return - - # Batch generate embeddings for nodes that need them - nodes_without_vectors = [n for n in nodes if n.vector is None and n.content] - if nodes_without_vectors: - nodes_with_vectors = await self.get_node_embeddings(nodes_without_vectors) - vector_map = {n.vector_id: n for n in nodes_with_vectors} - nodes_to_update = [vector_map.get(n.vector_id, n) if n.vector is None and n.content else n for n in nodes] - else: - nodes_to_update = nodes - - docs = [_vector_node_to_doc(n) for n in nodes_to_update] - self._collection.upsert(docs) - logger.info(f"Updated {len(nodes_to_update)} nodes in {self.collection_name}") - - async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode]: - """Fetch specific vector nodes from the collection by their IDs.""" - is_single = isinstance(vector_ids, str) - ids = [vector_ids] if is_single else vector_ids - - result_dict = self._collection.fetch(ids) - nodes = [_doc_to_vector_node(doc) for doc in result_dict.values()] - return nodes[0] if is_single and nodes else (nodes if not is_single else None) - - async def list( - self, - filters: dict | None = None, - limit: int | None = None, - sort_key: str | None = None, - reverse: bool = True, - ) -> list[VectorNode]: - """Retrieve vectors matching optional metadata filters. - - Uses zvec's ``query()`` without a vector query to list all documents. - Promoted metadata fields are filtered natively via zvec's ``filter`` - parameter; remaining filters are applied as post-filtering in Python. - - Args: - filters: Dictionary of filter conditions to match vectors. - limit: Maximum number of vectors to return. - sort_key: Key to sort the results by (in metadata). - reverse: If True, sort in descending order; otherwise ascending. - """ - # Split filters: native zvec filter vs Python post-filter - native_filter, post_filters = _build_zvec_filter(filters, self._promoted_fields_in_schema) - - # Determine fetch limit — zvec max topk is 1024 (will be lifted to 100,000 in zvec v0.3.2+) - _ZVEC_MAX_TOPK = 1024 - fetch_limit = min(limit or _ZVEC_MAX_TOPK, _ZVEC_MAX_TOPK) - if sort_key or post_filters: - fetch_limit = _ZVEC_MAX_TOPK # fetch max and sort/filter in Python - - results = self._collection.query( - topk=fetch_limit, - filter=native_filter, - include_vector=True, - ) - - nodes = [_doc_to_vector_node(doc) for doc in results] - - # Post-filter on non-promoted metadata fields - nodes = _apply_filters_post(nodes, post_filters) - - # Apply sorting if sort_key is provided - if sort_key: - - def _sort_key_func(node: VectorNode): - value = node.metadata.get(sort_key) - if value is None: - return float("-inf") if not reverse else float("inf") - return value - - nodes.sort(key=_sort_key_func, reverse=reverse) - - if limit is not None: - nodes = nodes[:limit] - - return nodes - - # ------------------------------------------------------------------ - # Helpers - # ------------------------------------------------------------------ - - def _detect_promoted_fields(self) -> None: - """Detect which promoted fields exist in the current collection's schema. - - Compares the set of promoted field names against the actual schema - and populates ``_promoted_fields_in_schema`` accordingly. Only fields - present in the schema can use native zvec filtering. - """ - if self._collection is None: - return - - try: - schema = self._collection.schema - existing_fields = {f.name for f in schema.fields} if schema.fields else set() - except Exception as e: - logger.warning(f"Failed to read collection schema: {e}") - existing_fields = set() - - self._promoted_fields_in_schema = set(_PROMOTED_FIELD_SPECS.keys()) & existing_fields - - missing = set(_PROMOTED_FIELD_SPECS.keys()) - existing_fields - if missing: - logger.info( - f"Promoted fields not in schema (will use post-filtering): {missing}", - ) - - def _ensure_numeric_promoted_columns(self) -> None: - """Add missing numeric promoted fields to existing collections. - - zvec's ``add_column`` only supports numeric types (INT64, FLOAT, etc.). - STRING fields cannot be added via ``add_column`` and must be defined - at collection creation time. For those, we fall back to post-filtering. - """ - if self._collection is None: - return - - missing = set(_PROMOTED_FIELD_SPECS.keys()) - self._promoted_fields_in_schema - if not missing: - return - - # Populate the DataType map if needed - if not _DATATYPE_MAP: - _DATATYPE_MAP.update( - { - "STRING": DataType.STRING, - "INT64": DataType.INT64, - }, - ) - - for field_name in missing: - type_str, _ = _PROMOTED_FIELD_SPECS[field_name] - # Only numeric types can be added via add_column - if type_str not in ("INT64", "INT32", "FLOAT", "DOUBLE"): - continue - try: - dt = _DATATYPE_MAP[type_str] - self._collection.add_column(FieldSchema(field_name, dt, nullable=True)) - self._promoted_fields_in_schema.add(field_name) - logger.info(f"Added promoted column '{field_name}' to existing collection") - except Exception as e: - logger.warning(f"Failed to add column '{field_name}': {e}") - - async def count(self) -> int: - """Return the total number of documents in the current collection.""" - stats = self._collection.stats - return stats.doc_count if stats else 0 - - async def reset(self): - """Reset the current collection by destroying and recreating it.""" - logger.warning(f"Resetting collection {self.collection_name}...") - await self.delete_collection(self.collection_name) - await self.create_collection(self.collection_name) - logger.info(f"Collection {self.collection_name} has been reset") diff --git a/reme4/enumeration/__init__.py b/reme/enumeration/__init__.py similarity index 100% rename from reme4/enumeration/__init__.py rename to reme/enumeration/__init__.py diff --git a/reme4/enumeration/chunk_enum.py b/reme/enumeration/chunk_enum.py similarity index 100% rename from reme4/enumeration/chunk_enum.py rename to reme/enumeration/chunk_enum.py diff --git a/reme4/enumeration/component_enum.py b/reme/enumeration/component_enum.py similarity index 100% rename from reme4/enumeration/component_enum.py rename to reme/enumeration/component_enum.py diff --git a/reme4/enumeration/link_scope_enum.py b/reme/enumeration/link_scope_enum.py similarity index 100% rename from reme4/enumeration/link_scope_enum.py rename to reme/enumeration/link_scope_enum.py diff --git a/reme/extension/__init__.py b/reme/extension/__init__.py deleted file mode 100644 index e3af3a08..00000000 --- a/reme/extension/__init__.py +++ /dev/null @@ -1,20 +0,0 @@ -"""Extension operations and tools.""" - -from . import procedural_memory -from .simple_chat import SimpleChat -from .stream_chat import StreamChat -from .test_op import TestOp -from .translate_ts import TranslateTs -from ..core.registry_factory import R - -__all__ = [ - "procedural_memory", - "SimpleChat", - "StreamChat", - "TestOp", - "TranslateTs", -] - -for name in __all__: - op_class = globals()[name] - R.ops.register(op_class) diff --git a/reme/extension/procedural_memory/__init__.py b/reme/extension/procedural_memory/__init__.py deleted file mode 100644 index 7a2cb783..00000000 --- a/reme/extension/procedural_memory/__init__.py +++ /dev/null @@ -1,47 +0,0 @@ -"""Procedural memory workflow.""" - -from ...core import R - -from .dump_memory import DumpMemory -from .load_memory import LoadMemory - -from .summary.trajectory_preprocess import TrajectoryPreprocess -from .summary.trajectory_segmentation import TrajectorySegmentation -from .summary.success_extraction import SuccessExtraction -from .summary.failure_extraction import FailureExtraction -from .summary.comparative_extraction import ComparativeExtraction -from .summary.memory_validation import MemoryValidation -from .summary.memory_deduplication import MemoryDeduplication -from .summary.memory_addition import MemoryAddition - -from .retrieve.build_query import BuildQuery -from .retrieve.memory_deletion import MemoryDeletion -from .retrieve.memory_retrieval import MemoryRetrieval -from .retrieve.merge_memory import MergeMemory -from .retrieve.rerank_memory import RerankMemory -from .retrieve.rewrite_memory import RewriteMemory -from .retrieve.update_memory_metadata import UpdateMemoryMetadata - -__all__ = [ - "DumpMemory", - "LoadMemory", - "TrajectoryPreprocess", - "TrajectorySegmentation", - "SuccessExtraction", - "FailureExtraction", - "ComparativeExtraction", - "MemoryValidation", - "MemoryDeduplication", - "MemoryAddition", - "BuildQuery", - "MemoryDeletion", - "MemoryRetrieval", - "MergeMemory", - "RerankMemory", - "RewriteMemory", - "UpdateMemoryMetadata", -] - -for name in __all__: - tool_class = globals()[name] - R.ops.register()(tool_class) diff --git a/reme/extension/procedural_memory/dump_memory.py b/reme/extension/procedural_memory/dump_memory.py deleted file mode 100644 index 8ca993d0..00000000 --- a/reme/extension/procedural_memory/dump_memory.py +++ /dev/null @@ -1,63 +0,0 @@ -"""Operation for dumping memories from vector store to JSONL file.""" - -import json -from pathlib import Path -from typing import List - -from loguru import logger - -from ...core.op import BaseOp -from ...core.schema.memory_node import MemoryNode -from ...core.schema.vector_node import VectorNode - - -class DumpMemory(BaseOp): - """Operation that dumps memories from vector store to a JSONL file. - - This operation retrieves all memories from the vector store, converts them - to MemoryNode objects, and writes them to a JSONL file (one JSON object - per line) for backup or export purposes. - """ - - async def execute(self): - """Execute the memory dump operation. - - Dumps all memories from the vector store to a JSONL file: - 1. Retrieves all VectorNodes from the vector store - 2. Converts them to MemoryNode objects - 3. Writes each MemoryNode as a JSON line to the output file - - Expected context attributes: - dump_file_path: Path to the output JSONL file. - - Sets context attributes: - dumped_count: Number of memories dumped to the file. - """ - # Support both dump_file_path and path for backward compatibility - dump_file_path: str = self.context.dump_file_path - if not dump_file_path: - logger.error("dump_file_path is required in context") - return - - file_path = Path(dump_file_path) - file_path.parent.mkdir(parents=True, exist_ok=True) - - # Retrieve all nodes from vector store - vector_nodes: List[VectorNode] = await self.vector_store.list() - logger.info(f"Retrieved {len(vector_nodes)} nodes from vector store") - - # Convert to MemoryNodes and write to JSONL file - dumped_count = 0 - with open(file_path, "w", encoding="utf-8") as f: - for node in vector_nodes: - try: - memory = MemoryNode.from_vector_node(node) - # Write as JSON line (one JSON object per line) - json_line = json.dumps(memory.model_dump(exclude_none=True), ensure_ascii=False) - f.write(json_line + "\n") - dumped_count += 1 - except Exception as e: - logger.warning(f"Failed to convert and dump node {node.vector_id}: {e}") - continue - - logger.info(f"Dumped {dumped_count} memories to {dump_file_path}") diff --git a/reme/extension/procedural_memory/load_memory.py b/reme/extension/procedural_memory/load_memory.py deleted file mode 100644 index b861765b..00000000 --- a/reme/extension/procedural_memory/load_memory.py +++ /dev/null @@ -1,82 +0,0 @@ -"""Operation for loading memories from JSONL file to vector store.""" - -import json -import asyncio -from pathlib import Path -from typing import List - -from loguru import logger - -from ...core.op import BaseOp -from ...core.schema.memory_node import MemoryNode -from ...core.schema.vector_node import VectorNode - - -class LoadMemory(BaseOp): - """Operation that loads memories from a JSONL file to vector store. - - This operation reads MemoryNode objects from a JSONL file (one JSON object - per line), converts them to VectorNode objects, and inserts them into the - vector store. - """ - - async def execute(self): - """Execute the memory load operation. - - Loads memories from a JSONL file to the vector store: - 1. Reads each line from the JSONL file - 2. Parses JSON and creates MemoryNode objects - 3. Converts MemoryNodes to VectorNodes - 4. Inserts them into the vector store - - Expected context attributes: - load_file_path: Path to the input JSONL file. - clear_existing: Optional. If True, clears existing memories before loading (default: False). - - Sets context attributes: - loaded_count: Number of memories loaded from the file. - """ - load_file_path: str = self.context.load_file_path - if not load_file_path: - logger.error("load_file_path is required in context") - return - - file_path = Path(load_file_path) - if not file_path.exists(): - logger.error(f"File not found: {load_file_path}") - return - - try: - # Attempt to retrieve the event loop associated with the current thread - loop = asyncio.get_running_loop() - print(f"Running event loop found: {loop}") - except RuntimeError: - # Start a new event loop to run the coroutine to completion - print("No running event loop found, starting a new one") - - clear_existing: bool = self.context.get("clear_existing", False) - if clear_existing: - await self.vector_store.delete_all() - logger.info("Cleared existing memories from vector store") - - # Read and parse JSONL file - memory_nodes: List[MemoryNode] = [] - with open(file_path, "r", encoding="utf-8") as f: - for line_num, line in enumerate(f, 1): - line = line.strip() - if not line: - continue - try: - data = json.loads(line) - memory = MemoryNode.model_validate(data) - memory_nodes.append(memory) - except Exception as e: - logger.warning(f"Failed to parse line {line_num} in {load_file_path}: {e}") - continue - logger.info(f"Parsed {len(memory_nodes)} memories from {load_file_path}") - - # Convert to VectorNodes and insert into vector store - if memory_nodes: - vector_nodes: List[VectorNode] = [memory.to_vector_node() for memory in memory_nodes] - await self.vector_store.insert(nodes=vector_nodes) - logger.info(f"Loaded {len(memory_nodes)} memories into vector store") diff --git a/reme/extension/procedural_memory/retrieve/__init__.py b/reme/extension/procedural_memory/retrieve/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/reme/extension/procedural_memory/retrieve/build_query.py b/reme/extension/procedural_memory/retrieve/build_query.py deleted file mode 100644 index dc121508..00000000 --- a/reme/extension/procedural_memory/retrieve/build_query.py +++ /dev/null @@ -1,58 +0,0 @@ -"""Query building operation module. - -This module provides functionality to build retrieval queries from either -explicit query strings or conversation messages, optionally using LLM to -generate optimized queries. -""" - -from loguru import logger - -from ....core.enumeration import Role -from ....core.op import BaseOp -from ....core.schema.message import Message -from ..utils import merge_messages_content - - -class BuildQuery(BaseOp): - """Build retrieval query from context or messages. - - This operation constructs a query string for memory retrieval. It can use - an explicit query from context, or generate one from conversation messages - using either LLM-based generation or simple message concatenation. - """ - - async def execute(self): - """Execute the query building operation. - - Builds a query string from either: - 1. An explicit query in the context - 2. Conversation messages (using LLM or simple concatenation) - - Stores the built query in context.query. - """ - if "query" in self.context: - query = self.context.query - - elif "messages" in self.context: - if self.context.get("enable_llm_build", True): - execution_process = merge_messages_content(self.context.messages) - prompt = self.prompt_format(prompt_name="query_build", execution_process=execution_process) - message = await self.llm.chat(messages=[Message(role=Role.USER, content=prompt)]) - query = message.content - - else: - context_parts = [] - message_summaries = [] - for message in self.context.messages[-3:]: # Last 3 messages - content = message.content[:200] + "..." if len(message.content) > 200 else message.content - message_summaries.append(f"- {message.role.value}: {content}") - if message_summaries: - context_parts.append("Recent messages:\n" + "\n".join(message_summaries)) - - query = "\n\n".join(context_parts) - - else: - raise RuntimeError("query or messages is required!") - - logger.info(f"build.query={query}") - self.context.query = query diff --git a/reme/extension/procedural_memory/retrieve/memory_deletion.py b/reme/extension/procedural_memory/retrieve/memory_deletion.py deleted file mode 100644 index 90315f74..00000000 --- a/reme/extension/procedural_memory/retrieve/memory_deletion.py +++ /dev/null @@ -1,54 +0,0 @@ -"""Operation for deleting memories from the vector store.""" - -import json -from typing import List - -from loguru import logger - -from ....core.op import BaseOp -from ....core.schema.vector_node import VectorNode - - -class MemoryDeletion(BaseOp): - """Operation that deletes memories from the vector store. - - This operation identifies memories to delete based on frequency and utility - thresholds, then deletes them. Memories with frequency >= freq_threshold - and utility/frequency ratio < utility_threshold are deleted. - """ - - async def execute(self): - """Execute the memory deletion operation. - - Identifies and deletes memories from the vector store: - 1. Lists all nodes from the vector store - 2. Identifies memories that meet deletion criteria based on thresholds - 3. Deletes identified memories from the vector store - 4. Stores deletion count in response.metadata["result"] - - The deletion criteria: - - Memory frequency must be >= freq_threshold - - Memory utility/frequency ratio must be < utility_threshold - - Expected context attributes: - freq_threshold: Minimum frequency threshold for consideration. - utility_threshold: Maximum utility/frequency ratio threshold. - """ - - # Step 1: Identify memories to delete based on thresholds - freq_threshold: int = self.context.freq_threshold - utility_threshold: float = self.context.utility_threshold - nodes: List[VectorNode] = await self.vector_store.list() - - deleted_memory_ids = [] - for node in nodes: - freq = node.metadata.get("freq", 0) - utility = node.metadata.get("utility", 0) - if freq >= freq_threshold: - if freq > 0 and utility * 1.0 / freq < utility_threshold: - deleted_memory_ids.append(node.vector_id) - - # Step 2: Execute deletion if there are any IDs to delete - if deleted_memory_ids: - await self.vector_store.delete(vector_ids=deleted_memory_ids) - logger.info(f"Deleted {len(deleted_memory_ids)} memories: {json.dumps(deleted_memory_ids, indent=2)}") diff --git a/reme/extension/procedural_memory/retrieve/memory_retrieval.py b/reme/extension/procedural_memory/retrieve/memory_retrieval.py deleted file mode 100644 index c5c9b5b8..00000000 --- a/reme/extension/procedural_memory/retrieve/memory_retrieval.py +++ /dev/null @@ -1,68 +0,0 @@ -"""Operation for recalling memories from the vector store based on a query.""" - -from typing import List - -from loguru import logger - -from ....core.op import BaseOp -from ....core.schema.memory_node import MemoryNode -from ....core.schema.vector_node import VectorNode - - -class MemoryRetrieval(BaseOp): - """Operation that retrieves relevant memories from the vector store. - - This operation performs a semantic search on the vector store to find - memories relevant to a given query. It supports optional score filtering - and deduplication based on memory content. - """ - - async def execute(self): - """Execute the memory recall operation. - - Performs a semantic search in the vector store using the provided query, - retrieves the top-k most relevant memories, and optionally filters them - by a score threshold. Duplicate memories (based on content) are removed. - - Expected context attributes: - query: The search query string. - top_k: Number of top results to retrieve (default: 3). - - Expected context attributes (optional): - threshold_score: Optional minimum score threshold for filtering. - - Sets response.metadata: - memory_list: List of retrieved MemoryNode objects. - """ - top_k: int = self.context.get("top_k", 5) - - query: str = self.context.get("query", "") - assert query, "query should be not empty!" - - # Perform semantic search - nodes: List[VectorNode] = await self.vector_store.search( - query=query, - limit=top_k, - filters=None, - ) - - # Convert VectorNodes to MemoryNodes and deduplicate by content - memory_list: List[MemoryNode] = [] - memory_content_set: set[str] = set() # for deduplication - for node in nodes: - try: - memory = MemoryNode.from_vector_node(node) - if memory.content not in memory_content_set: - memory_list.append(memory) - memory_content_set.add(memory.content) - except Exception as e: - logger.warning(f"Failed to convert VectorNode to MemoryNode: {e}") - continue - logger.info(f"Retrieved memory.size={len(memory_list)}") - - threshold_score: float | None = self.context.get("threshold_score", None) - if threshold_score is not None: - memory_list = [mem for mem in memory_list if mem.score >= threshold_score] - logger.info(f"After threshold filter: {len(memory_list)} memories retained") - - self.context.response.metadata["memory_list"] = memory_list diff --git a/reme/extension/procedural_memory/retrieve/merge_memory.py b/reme/extension/procedural_memory/retrieve/merge_memory.py deleted file mode 100644 index c54f5806..00000000 --- a/reme/extension/procedural_memory/retrieve/merge_memory.py +++ /dev/null @@ -1,45 +0,0 @@ -"""Memory merging operation module. - -This module provides functionality to merge multiple retrieved memories -into a single formatted context string for use in LLM responses. -""" - -from typing import List - -from loguru import logger - -from ....core.op import BaseOp -from ....core.schema.memory_node import MemoryNode - - -class MergeMemory(BaseOp): - """Merge multiple memories into a single formatted context. - - This operation takes a list of retrieved memories and formats them into - a single context string that can be used to guide LLM responses. It includes - instructions for the LLM to consider the helpful parts from these memories. - """ - - async def execute(self): - """Execute the memory merging operation. - - Merges memories from context metadata into a formatted string with - instructions for the LLM. Stores the merged result in response.answer. - """ - memory_list: List[MemoryNode] = self.context.response.metadata["memory_list"] - - if not memory_list: - return - - content_collector = ["Previous Memory"] - for memory in memory_list: - if not memory.content: - continue - - content_collector.append(f"- {memory.when_to_use} {memory.content}\n") - content_collector.append( - "Please consider the helpful parts from these in answering the question, " - "to make the response more comprehensive and substantial.", - ) - self.context.response.answer = "\n".join(content_collector) - logger.info(f"response.answer={self.context.response.answer}") diff --git a/reme/extension/procedural_memory/retrieve/rerank_memory.py b/reme/extension/procedural_memory/retrieve/rerank_memory.py deleted file mode 100644 index 12d73815..00000000 --- a/reme/extension/procedural_memory/retrieve/rerank_memory.py +++ /dev/null @@ -1,186 +0,0 @@ -"""Memory reranking operation module. - -This module provides functionality to rerank and filter retrieved memories -using LLM-based reranking and score-based filtering to select the most relevant -memories for the current task. -""" - -import json -import re -from typing import List - -from loguru import logger - -from ....core.enumeration import Role -from ....core.op import BaseOp -from ....core.schema.message import Message -from ....core.schema.memory_node import MemoryNode - - -class RerankMemory(BaseOp): - """Rerank and filter recalled experiences using LLM and score-based filtering. - - This operation takes recalled memories and applies multiple filtering and - ranking strategies to select the most relevant memories for the current task. - It supports LLM-based reranking and score-based filtering. - """ - - async def execute(self): - """Execute the memory reranking operation. - - Applies LLM-based reranking (optional) and score-based filtering (optional) - to rerank retrieved memories. Stores the reranked results - in the context response metadata. - """ - memory_list: List[MemoryNode] = self.context.response.metadata["memory_list"] - retrieval_query: str = self.context.query - enable_llm_rerank = self.context.get("enable_llm_rerank", False) - enable_score_filter = self.context.get("enable_score_filter", False) - min_score_threshold = self.context.get("min_score_threshold", 0.3) - - if not memory_list: - logger.info("No recalled memory_list to rerank") - return - - logger.info(f"Reranking {len(memory_list)} memories") - - # Step 1: LLM reranking (optional) - if enable_llm_rerank: - memory_list = await self._llm_rerank(retrieval_query, memory_list) - logger.info(f"After LLM reranking: {len(memory_list)} memories") - - # Step 2: Score-based filtering (optional) - if enable_score_filter: - memory_list = self._score_based_filter(memory_list, min_score_threshold) - logger.info(f"After score filtering: {len(memory_list)} memories") - - # Store results in context - self.context.response.metadata["memory_list"] = memory_list - - async def _llm_rerank(self, query: str, candidates: List[MemoryNode]) -> List[MemoryNode]: - """LLM-based reranking of candidate experiences. - - Args: - query: The retrieval query used to rank candidates. - candidates: List of memory candidates to rerank. - - Returns: - List of memories reranked by relevance to the query. - """ - if not candidates: - return candidates - - # Format candidates for LLM evaluation - candidates_text = self._format_candidates_for_rerank(candidates) - - prompt = self.prompt_format( - prompt_name="memory_rerank_prompt", - query=query, - candidates=candidates_text, - num_candidates=len(candidates), - ) - - response = await self.llm.chat(messages=[Message(role=Role.USER, content=prompt)]) - - # Parse reranking results - reranked_indices = self._parse_rerank_response(response.content) - - # Reorder candidates based on LLM ranking - if reranked_indices: - reranked_candidates = [] - for idx in reranked_indices: - if 0 <= idx < len(candidates): - reranked_candidates.append(candidates[idx]) - - # Add any remaining candidates that weren't explicitly ranked - ranked_indices_set = set(reranked_indices) - for i, candidate in enumerate(candidates): - if i not in ranked_indices_set: - reranked_candidates.append(candidate) - - return reranked_candidates - - return candidates - - @staticmethod - def _score_based_filter(memories: List[MemoryNode], min_score: float) -> List[MemoryNode]: - """Filter memories based on quality scores. - - Args: - memories: List of memories to filter. - min_score: Minimum combined score threshold for filtering. - - Returns: - List of memories that meet the minimum score threshold. - """ - filtered_memories = [] - - for memory in memories: - # Get confidence score from metadata - confidence = memory.metadata.get("confidence", 0.5) - validation_score = memory.score or 0.5 - - # Calculate combined score - combined_score = (confidence + validation_score) / 2 - - if combined_score >= min_score: - filtered_memories.append(memory) - else: - logger.debug(f"Filtered out memory with score {combined_score:.2f}") - - logger.info(f"Score filtering: {len(filtered_memories)}/{len(memories)} memories retained") - return filtered_memories - - @staticmethod - def _format_candidates_for_rerank(candidates: List[MemoryNode]) -> str: - """Format candidates for LLM reranking. - - Args: - candidates: List of memory candidates to format. - - Returns: - Formatted string representation of candidates for LLM evaluation. - """ - formatted_candidates = [] - - for i, candidate in enumerate(candidates): - condition = candidate.when_to_use - content = candidate.content - - candidate_text = f"Candidate {i}:\n" - candidate_text += f"Condition: {condition}\n" - candidate_text += f"Experience: {content}\n" - - formatted_candidates.append(candidate_text) - - return "\n---\n".join(formatted_candidates) - - @staticmethod - def _parse_rerank_response(response: str) -> List[int]: - """Parse LLM reranking response to extract ranked indices. - - Args: - response: The LLM response containing ranked indices. - - Returns: - List of indices representing the reranked order. - """ - try: - # Try to extract JSON format - json_pattern = r"```json\s*([\s\S]*?)\s*```" - json_blocks = re.findall(json_pattern, response) - - if json_blocks: - parsed = json.loads(json_blocks[0]) - if isinstance(parsed, dict) and "ranked_indices" in parsed: - return parsed["ranked_indices"] - elif isinstance(parsed, list): - return parsed - - # Try to extract numbers from text - numbers = re.findall(r"\b\d+\b", response) - return [int(num) for num in numbers if int(num) < 100] # Reasonable upper bound - - except Exception as e: - logger.error(f"Error parsing rerank response: {e}") - return [] diff --git a/reme/extension/procedural_memory/retrieve/rewrite_memory.py b/reme/extension/procedural_memory/retrieve/rewrite_memory.py deleted file mode 100644 index a60dc7e0..00000000 --- a/reme/extension/procedural_memory/retrieve/rewrite_memory.py +++ /dev/null @@ -1,201 +0,0 @@ -"""Memory rewriting operation module. - -This module provides functionality to rewrite and format retrieved memories -into context messages that can be used by LLMs for task completion. -""" - -import json -import re -from typing import List - -from loguru import logger - -from ....core.enumeration import Role -from ....core.op import BaseOp -from ....core.schema.message import Message -from ....core.schema.memory_node import MemoryNode - - -class RewriteMemory(BaseOp): - """Generate and rewrite context messages from reranked experiences. - - This operation takes reranked memories and formats them into context messages - that can be used by LLMs. It optionally uses LLM-based rewriting to make - the context more relevant and actionable for the current task. - """ - - async def execute(self): - """Execute the memory rewrite operation. - - Retrieves memories from context metadata, formats them, and optionally - rewrites them using LLM to make them more relevant for the current query. - Stores the rewritten context in the response answer field. - """ - memory_list: List[MemoryNode] = self.context.response.metadata["memory_list"] - query: str = self.context.query - messages: List[Message] = [Message(**x) if isinstance(x, dict) else x for x in self.context.get("messages", [])] - - if not memory_list: - logger.info("No reranked memories to rewrite") - self.context.response.answer = "" - return - - logger.info(f"Generating context from {len(memory_list)} memories") - - # Generate initial context message - rewritten_memory = await self._generate_context_message(query, messages, memory_list) - - # Store results in context - self.context.response.answer = rewritten_memory - self.context.response.metadata["memory_list"] = [memory.model_dump() for memory in memory_list] - - async def _generate_context_message(self, query: str, messages: List[Message], memories: List[MemoryNode]) -> str: - """Generate context message from retrieved memories. - - Args: - query: The current query string. - messages: List of conversation messages for context. - memories: List of retrieved memories to format. - - Returns: - Formatted context string, optionally rewritten by LLM. - """ - if not memories: - return "" - - try: - logger.info("memories") - # Format retrieved memories - formatted_memories = self._format_memories_for_context(memories) - - if self.context.get("enable_llm_rewrite", False): - context_content = await self._rewrite_context(query, formatted_memories, messages) - else: - context_content = formatted_memories - - return context_content - - except Exception as e: - logger.error(f"Error generating context message: {e}") - return self._format_memories_for_context(memories) - - async def _rewrite_context(self, query: str, context_content: str, messages: List[Message]) -> str: - """LLM-based context rewriting to make experiences more relevant and actionable. - - Args: - query: The current query string. - context_content: The formatted context content to rewrite. - messages: List of conversation messages for additional context. - - Returns: - Rewritten context string optimized for the current task. - """ - if not context_content: - return context_content - - try: - # Extract current context - current_context = self._extract_context(messages) - - prompt = self.prompt_format( - prompt_name="memory_rewrite_prompt", - current_query=query, - current_context=current_context, - original_context=context_content, - ) - - response = await self.llm.chat(messages=[Message(role=Role.USER, content=prompt)]) - - # Extract rewritten context - rewritten_context = self._parse_json_response(response.content, "rewritten_context") - - if rewritten_context and rewritten_context.strip(): - logger.info("Context successfully rewritten for current task") - return rewritten_context.strip() - - return context_content - - except Exception as e: - logger.error(f"Error in context rewriting: {e}") - return context_content - - @staticmethod - def _format_memories_for_context(memories: List[MemoryNode]) -> str: - """Format memories for context generation. - - Args: - memories: List of memories to format. - - Returns: - Formatted string containing all memories with their conditions and content. - """ - formatted_memories = [] - - for i, memory in enumerate(memories, 1): - condition = memory.when_to_use - memory_content = memory.content - memory_text = f"Memory {i} :\n When to use: {condition}\n Content: {memory_content}\n" - - formatted_memories.append(memory_text) - - return "\n".join(formatted_memories) - - @staticmethod - def _extract_context(messages: List[Message]) -> str: - """Extract relevant context from messages. - - Args: - messages: List of conversation messages. - - Returns: - Formatted string containing recent conversation context. - """ - if not messages: - return "" - - context_parts = [] - - # Add recent messages if available - recent_messages = messages[-3:] # Last 3 messages - message_summaries = [] - for message in recent_messages: - content = message.content[:300] + "..." if len(message.content) > 300 else message.content - message_summaries.append(f"- {message.role.value}: {content}") - - if message_summaries: - context_parts.append("Recent conversation:\n" + "\n".join(message_summaries)) - - return "\n\n".join(context_parts) - - @staticmethod - def _parse_json_response(response: str, key: str) -> str: - """Parse JSON response to extract specific key. - - Args: - response: The response string that may contain JSON. - key: The key to extract from the JSON object. - - Returns: - The value associated with the key, or the response string if parsing fails. - """ - try: - # Try to extract JSON blocks - json_pattern = r"```json\s*([\s\S]*?)\s*```" - json_blocks = re.findall(json_pattern, response) - - if json_blocks: - parsed = json.loads(json_blocks[0]) - if isinstance(parsed, dict) and key in parsed: - return parsed[key] - - # Fallback: try to parse the entire response as JSON - parsed = json.loads(response) - if isinstance(parsed, dict) and key in parsed: - return parsed[key] - - except json.JSONDecodeError: - logger.warning(f"Failed to parse JSON response for key '{key}', using raw response") - # If JSON parsing fails, return the response as-is for fallback - return response.strip() - - return "" diff --git a/reme/extension/procedural_memory/retrieve/update_memory_metadata.py b/reme/extension/procedural_memory/retrieve/update_memory_metadata.py deleted file mode 100644 index 14d7e57a..00000000 --- a/reme/extension/procedural_memory/retrieve/update_memory_metadata.py +++ /dev/null @@ -1,51 +0,0 @@ -"""Operation for updating memory metadata (frequency and utility). - -This module provides a unified operation to update frequency counters and -optionally utility scores for recalled memories, directly updating the -vector store. -""" - -from typing import List - -from loguru import logger - -from ....core.op import BaseOp -from ....core.schema.memory_node import MemoryNode -from ....core.schema.vector_node import VectorNode - - -class UpdateMemoryMetadata(BaseOp): - """Update memory metadata: frequency and optionally utility. - - This operation (1) increments each memory's frequency counter; - (2) optionally increments utility when update_utility is True; - (3) directly updates the VectorNode in the vector store using the update method. - - Expected context attributes: - memory_list: List of MemoryNode objects to update (already loaded from - previous operations like rerank_memory). - update_utility: Boolean flag. If True, also increment utility for each memory. - """ - - async def execute(self): - """Run frequency update, optional utility update, and directly update vector store.""" - memory_list: List[MemoryNode] = [MemoryNode(**node) for node in self.context.memory_list] - update_utility = self.context.update_utility - - if not memory_list: - logger.info("No memories to update metadata") - return - - updated_nodes: List[VectorNode] = [] - for memory in memory_list: - meta = memory.metadata - meta["freq"] = meta.get("freq", 0) + 1 - if update_utility: - meta["utility"] = meta.get("utility", 0) + 1 - memory.metadata = meta - vector_node = memory.to_vector_node() - updated_nodes.append(vector_node) - - if updated_nodes: - await self.vector_store.update(nodes=updated_nodes) - logger.info(f"Updated metadata for {len(updated_nodes)} memories in vector store") diff --git a/reme/extension/procedural_memory/summary/__init__.py b/reme/extension/procedural_memory/summary/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/reme/extension/procedural_memory/summary/comparative_extraction.py b/reme/extension/procedural_memory/summary/comparative_extraction.py deleted file mode 100644 index 4e728de3..00000000 --- a/reme/extension/procedural_memory/summary/comparative_extraction.py +++ /dev/null @@ -1,274 +0,0 @@ -"""Comparative extraction operation for task memory generation. - -This module provides operations to extract comparative task memories by comparing -different trajectories with varying scores or success/failure outcomes. -""" - -from typing import List, Tuple, Optional - -from loguru import logger - -from ....core.enumeration import MemoryType, Role -from ....core.op import BaseOp -from ....core.schema.memory_node import MemoryNode -from ....core.schema.message import Message, Trajectory -from ..utils import ( - merge_messages_content, - parse_json_experience_response, -) - - -class ComparativeExtraction(BaseOp): - """Extract comparative task memories by comparing different scoring trajectories. - - This operation performs two types of comparisons: - 1. Soft comparison: Compares highest vs lowest scoring trajectories - 2. Hard comparison: Compares similar success vs failure step sequences - - The extracted memories help identify what makes some trajectories more successful - than others. - """ - - async def 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", []) - failure_trajectories: List[Trajectory] = self.context.get("failure_trajectories", []) - - comparative_task_memories = [] - - # Soft comparison: highest score vs lowest score - if self.context.get("enable_soft_comparison", True) and len(all_trajectories) >= 2: - highest_traj, lowest_traj = self._find_highest_lowest_scoring_trajectories(all_trajectories) - if highest_traj and lowest_traj and highest_traj.score > lowest_traj.score: - logger.info( - f"Extracting soft comparative task memories: " - f"highest ({highest_traj.score:.2f}) vs lowest ({lowest_traj.score:.2f})", - ) - 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) - if self.context.get("enable_similarity_comparison", False) and success_trajectories and failure_trajectories: - similar_pairs = await self._find_similar_step_sequences(success_trajectories, failure_trajectories) - 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 = await self._extract_hard_comparative_task_memory( - success_steps, - failure_steps, - similarity_score, - ) - comparative_task_memories.extend(hard_task_memories) - - logger.info(f"Extracted {len(comparative_task_memories)} comparative task memories") - - # Add task memories to context - self.context.comparative_task_memories = comparative_task_memories - - @staticmethod - def _find_highest_lowest_scoring_trajectories(trajectories: List[Trajectory]) -> Tuple[ - Optional[Trajectory], - Optional[Trajectory], - ]: - """Find the highest and lowest scoring trajectories""" - if len(trajectories) < 2: - return None, None - - # Filter trajectories with valid scores - valid_trajectories = [traj for traj in trajectories if traj.score is not None] - - if len(valid_trajectories) < 2: - logger.warning("Not enough trajectories with valid scores for comparison") - return None, None - - # Sort by score - sorted_trajectories = sorted(valid_trajectories, key=lambda x: x.score, reverse=True) - - highest_traj = sorted_trajectories[0] - lowest_traj = sorted_trajectories[-1] - - return highest_traj, lowest_traj - - @staticmethod - def _get_trajectory_score(trajectory: Trajectory) -> Optional[float]: - """Get trajectory score""" - return trajectory.score - - async def _extract_soft_comparative_task_memory( - self, - higher_traj: Trajectory, - lower_traj: Trajectory, - ) -> List[MemoryNode]: - """Extract soft comparative task memory (high score vs low score)""" - higher_steps = self._get_trajectory_steps(higher_traj) - lower_steps = self._get_trajectory_steps(lower_traj) - higher_score = self._get_trajectory_score(higher_traj) - lower_score = self._get_trajectory_score(lower_traj) - - prompt = self.prompt_format( - prompt_name="soft_comparative_step_task_memory_prompt", - higher_steps=merge_messages_content(higher_steps), - lower_steps=merge_messages_content(lower_steps), - higher_score=f"{higher_score:.2f}", - lower_score=f"{lower_score:.2f}", - ) - - def parse_task_memories(message: Message) -> List[MemoryNode]: - task_memories_data = parse_json_experience_response(message.content) - task_memories = [] - - for tm_data in task_memories_data: - task_memory = MemoryNode( - memory_type=MemoryType.PROCEDURAL, - when_to_use=tm_data.get("when_to_use", tm_data.get("condition", "")), - content=tm_data.get("experience", ""), - author=getattr(self.llm, "model_name", "system"), - metadata=tm_data, - ) - task_memories.append(task_memory) - - return task_memories - - return await self.llm.chat( - messages=[Message(role=Role.USER, content=prompt)], - callback_fn=parse_task_memories, - ) - - async def _extract_hard_comparative_task_memory( - self, - success_steps: List[Message], - failure_steps: List[Message], - similarity_score: float, - ) -> List[MemoryNode]: - """Extract hard comparative task memory (success vs failure)""" - prompt = self.prompt_format( - prompt_name="hard_comparative_step_task_memory_prompt", - success_steps=merge_messages_content(success_steps), - failure_steps=merge_messages_content(failure_steps), - similarity_score=similarity_score, - ) - - def parse_task_memories(message: Message) -> List[MemoryNode]: - task_memories_data = parse_json_experience_response(message.content) - task_memories = [] - - for tm_data in task_memories_data: - task_memory = MemoryNode( - memory_type=MemoryType.PROCEDURAL, - when_to_use=tm_data.get("when_to_use", tm_data.get("condition", "")), - content=tm_data.get("experience", ""), - author=getattr(self.llm, "model_name", "system"), - metadata=tm_data, - ) - task_memories.append(task_memory) - - return task_memories - - return await self.llm.chat( - messages=[Message(role=Role.USER, content=prompt)], - callback_fn=parse_task_memories, - ) - - @staticmethod - def _get_trajectory_steps(trajectory: Trajectory) -> List[Message]: - """Get trajectory steps, prioritizing segmented steps""" - if hasattr(trajectory, "segments") and trajectory.segments: - # If there are segments, merge all segments - all_steps = [] - for segment in trajectory.segments: - all_steps.extend(segment) - return all_steps - else: - return trajectory.messages - - async def _find_similar_step_sequences( - self, - success_trajectories: List[Trajectory], - failure_trajectories: List[Trajectory], - ) -> List[Tuple[List[Message], List[Message], float]]: - """Find similar step sequences for comparison""" - try: - similar_pairs = [] - - # Get step sequences - success_step_sequences = [] - for traj in success_trajectories: - if hasattr(traj.metadata, "segments") and traj.metadata["segments"]: - success_step_sequences.extend(traj.metadata["segments"]) - else: - success_step_sequences.append(traj.messages) - - failure_step_sequences = [] - for traj in failure_trajectories: - if hasattr(traj.metadata, "segments") and traj.metadata["segments"]: - failure_step_sequences.extend(traj.metadata["segments"]) - else: - failure_step_sequences.append(traj.messages) - - # Limit comparison count to avoid computational overload - max_sequences = self.context.get("max_similarity_sequences", 5) - success_step_sequences = success_step_sequences[:max_sequences] - failure_step_sequences = failure_step_sequences[:max_sequences] - - if not success_step_sequences or not failure_step_sequences: - return [] - - # Generate text representation for embedding - success_texts = [merge_messages_content(seq) for seq in success_step_sequences] - failure_texts = [merge_messages_content(seq) for seq in failure_step_sequences] - - # Get embedding vectors - if ( - hasattr(self, "vector_store") - and self.vector_store - and hasattr( - self.vector_store, - "embedding_model", - ) - ): - success_embeddings = await self.vector_store.get_embeddings(success_texts) - failure_embeddings = await self.vector_store.get_embeddings(failure_texts) - - # Calculate similarity and find most similar pairs - similarity_threshold = self.context.get("similarity_threshold", 0.5) - - for i, s_emb in enumerate(success_embeddings): - for j, f_emb in enumerate(failure_embeddings): - similarity = self._calculate_cosine_similarity(s_emb, f_emb) - - if similarity > similarity_threshold: - similar_pairs.append( - ( - success_step_sequences[i], - failure_step_sequences[j], - similarity, - ), - ) - - # Return top most similar pairs - max_pairs = self.context.get("max_similarity_pairs", 3) - return sorted(similar_pairs, key=lambda x: x[2], reverse=True)[:max_pairs] - - except Exception as e: - logger.error(f"Error finding similar step sequences: {e}") - - return [] - - @staticmethod - def _calculate_cosine_similarity(embedding1: List[float], embedding2: List[float]) -> float: - """Calculate cosine similarity""" - import numpy as np - - vec1 = np.array(embedding1) - vec2 = np.array(embedding2) - - # Calculate cosine similarity - dot_product = np.dot(vec1, vec2) - norm1 = np.linalg.norm(vec1) - norm2 = np.linalg.norm(vec2) - - if norm1 == 0 or norm2 == 0: - return 0.0 - - return dot_product / (norm1 * norm2) diff --git a/reme/extension/procedural_memory/summary/failure_extraction.py b/reme/extension/procedural_memory/summary/failure_extraction.py deleted file mode 100644 index 42e682d8..00000000 --- a/reme/extension/procedural_memory/summary/failure_extraction.py +++ /dev/null @@ -1,95 +0,0 @@ -"""Failure extraction operation for task memory generation. - -This module provides operations to extract task memories from failed trajectories, -identifying mistakes, pitfalls, and lessons learned from failures. -""" - -from typing import List - -from loguru import logger - -from ....core.enumeration import MemoryType, Role -from ....core.op import BaseOp -from ....core.schema.memory_node import MemoryNode -from ....core.schema.message import Message, Trajectory -from ..utils import ( - get_trajectory_context, - merge_messages_content, - parse_json_experience_response, -) - - -class FailureExtraction(BaseOp): - """Extract task memories from failed trajectories. - - This operation analyzes failed trajectories (or their segments) to extract - lessons learned, common mistakes, and anti-patterns that should be avoided - in similar future tasks. - """ - - async def execute(self): - """Extract task memories from failed trajectories""" - failure_trajectories: List[Trajectory] = self.context.get("failure_trajectories", []) - - if not failure_trajectories: - logger.info("No failure trajectories found for extraction") - return - - logger.info(f"Extracting task memories from {len(failure_trajectories)} failed trajectories") - - failure_task_memories = [] - - # Process trajectories - for trajectory in failure_trajectories: - if "segments" in trajectory.metadata: - # Process segmented step sequences - for segment in trajectory.metadata["segments"]: - task_memories = await self._extract_failure_task_memory_from_steps(segment, trajectory) - failure_task_memories.extend(task_memories) - else: - # Process entire 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") - - # Add task memories to context - self.context.failure_task_memories = failure_task_memories - - async def _extract_failure_task_memory_from_steps( - self, - steps: List[Message], - trajectory: Trajectory, - ) -> List[MemoryNode]: - """Extract task memory from failed step sequences""" - step_content = merge_messages_content(steps) - context = get_trajectory_context(trajectory, steps) - - prompt = self.prompt_format( - prompt_name="failure_step_task_memory_prompt", - query=trajectory.metadata.get("query", ""), - step_sequence=step_content, - context=context, - outcome="failed", - ) - - def parse_task_memories(message: Message) -> List[MemoryNode]: - task_memories_data = parse_json_experience_response(message.content) - task_memories = [] - - for tm_data in task_memories_data: - task_memory = MemoryNode( - memory_type=MemoryType.PROCEDURAL, - when_to_use=tm_data.get("when_to_use", tm_data.get("condition", "")), - content=tm_data.get("experience", ""), - author=getattr(self.llm, "model_name", "system"), - metadata=tm_data, - ) - task_memories.append(task_memory) - - return task_memories - - return await self.llm.chat( - messages=[Message(role=Role.USER, content=prompt)], - callback_fn=parse_task_memories, - ) diff --git a/reme/extension/procedural_memory/summary/memory_addition.py b/reme/extension/procedural_memory/summary/memory_addition.py deleted file mode 100644 index ed56fdc9..00000000 --- a/reme/extension/procedural_memory/summary/memory_addition.py +++ /dev/null @@ -1,32 +0,0 @@ -"""Operation for adding memories to the vector store.""" - -from typing import List -from loguru import logger -from ....core.op import BaseOp -from ....core.schema.vector_node import VectorNode -from ....core.schema.memory_node import MemoryNode - - -class MemoryAddition(BaseOp): - """Operation that adds new or updated memories to the vector store. - - This operation inserts memories into the vector store. It reads the list - of memories to insert from response.metadata and performs the actual - database insertion operations. - """ - - async def execute(self): - """Execute the memory insertion operation. - - Inserts new or updated memories into the vector store: - 1. Reads memory_list from context (can be dicts or MemoryNode) - 2. Converts raw items to MemoryNode objects - 3. Converts MemoryNode objects to VectorNode objects - 4. Inserts them into the vector store - """ - raw_memory_list = self.context.memory_list - insert_memory_list: List[MemoryNode] = [MemoryNode(**x) if isinstance(x, dict) else x for x in raw_memory_list] - if insert_memory_list: - insert_nodes: List[VectorNode] = [x.to_vector_node() for x in insert_memory_list] - await self.vector_store.insert(nodes=insert_nodes) - logger.info(f"insert insert_node.size={len(insert_nodes)}") diff --git a/reme/extension/procedural_memory/summary/memory_deduplication.py b/reme/extension/procedural_memory/summary/memory_deduplication.py deleted file mode 100644 index 7d71051e..00000000 --- a/reme/extension/procedural_memory/summary/memory_deduplication.py +++ /dev/null @@ -1,183 +0,0 @@ -"""Memory deduplication operation for task memory management. - -This module provides operations to remove duplicate or highly similar task -memories by comparing embeddings and calculating similarity scores. -""" - -from typing import List - -from loguru import logger - -from ....core.op import BaseOp -from ....core.schema.memory_node import MemoryNode - - -class MemoryDeduplication(BaseOp): - """Remove duplicate task memories using embedding similarity. - - This operation identifies and removes duplicate or highly similar task - memories by comparing their embeddings against both existing memories - in the vector store and other memories in the current batch. - """ - - async def execute(self): - """Remove duplicate task memories""" - # Get task memories to deduplicate - task_memories: List[MemoryNode] = self.context.response.metadata.get("memory_list", []) - - if not task_memories: - logger.info("No task memories found for deduplication") - return - - logger.info(f"Starting deduplication for {len(task_memories)} task memories") - - # Perform deduplication - deduplicated_task_memories = await self._deduplicate_task_memories(task_memories) - - logger.info( - f"Deduplication complete: {len(deduplicated_task_memories)} deduplicated " - f"task memories out of {len(task_memories)}", - ) - - # Update context - self.context.response.metadata["memory_list"] = deduplicated_task_memories - - async def _deduplicate_task_memories(self, task_memories: List[MemoryNode]) -> List[MemoryNode]: - """Remove duplicate task memories""" - if not task_memories: - return task_memories - - similarity_threshold = self.context.get("similarity_threshold", 0.5) - - unique_task_memories = [] - - # Get existing task memory embeddings - existing_embeddings = await self._get_existing_task_memory_embeddings() - - for task_memory in task_memories: - # Generate embedding for current task memory - current_embedding = await self._get_task_memory_embedding(task_memory) - - if current_embedding is None: - logger.warning(f"Failed to generate embedding for task memory: {str(task_memory.when_to_use)[:50]}...") - continue - - # Check similarity with existing task memories - if self._is_similar_to_existing_task_memories(current_embedding, existing_embeddings, similarity_threshold): - logger.debug(f"Skipping similar task memory: {str(task_memory.when_to_use)[:50]}...") - continue - - # Check similarity with current batch task memories - if await self._is_similar_to_current_task_memories( - current_embedding, - unique_task_memories, - similarity_threshold, - ): - logger.debug(f"Skipping duplicate in current batch: {str(task_memory.when_to_use)[:50]}...") - continue - - # Add to unique task memories list - unique_task_memories.append(task_memory) - logger.debug(f"Added unique task memory: {str(task_memory.when_to_use)[:50]}...") - - return unique_task_memories - - async def _get_existing_task_memory_embeddings(self) -> List[List[float]]: - """Get embeddings of existing task memories""" - try: - if not hasattr(self, "vector_store") or not self.vector_store: - return [] - - # List all existing task memory nodes - existing_nodes = await self.vector_store.list( - filters=None, # No filters to get all nodes - limit=self.context.get("max_existing_task_memories", 1000), - ) - - # Extract embeddings - existing_embeddings = [] - for node in existing_nodes: - if node.vector: - existing_embeddings.append(node.vector) - - logger.debug( - f"Retrieved {len(existing_embeddings)} existing task memory embeddings", - ) - return existing_embeddings - - except Exception as e: - logger.warning(f"Failed to retrieve existing task memory embeddings: {e}") - return [] - - async def _get_task_memory_embedding(self, task_memory: MemoryNode) -> List[float] | None: - """Generate embedding for task memory""" - try: - - # Combine task memory description and content for embedding - text_for_embedding = f"{task_memory.when_to_use} {task_memory.content}" - embeddings = await self.vector_store.get_embeddings([text_for_embedding]) - - if embeddings and len(embeddings) > 0: - return embeddings[0] - else: - logger.warning("Empty embedding generated for task memory") - return None - - except Exception as e: - logger.error(f"Error generating embedding for task memory: {e}") - return None - - def _is_similar_to_existing_task_memories( - self, - current_embedding: List[float], - existing_embeddings: List[List[float]], - threshold: float, - ) -> bool: - """Check if current embedding is similar to existing embeddings""" - for existing_embedding in existing_embeddings: - similarity = self._calculate_cosine_similarity(current_embedding, existing_embedding) - if similarity > threshold: - logger.debug(f"Found similar existing task memory with similarity: {similarity:.3f}") - return True - return False - - async def _is_similar_to_current_task_memories( - self, - current_embedding: List[float], - current_task_memories: List[MemoryNode], - threshold: float, - ) -> bool: - """Check if current embedding is similar to other memories in current batch.""" - for existing_task_memory in current_task_memories: - existing_embedding = await self._get_task_memory_embedding(existing_task_memory) - if existing_embedding is None: - continue - - similarity = self._calculate_cosine_similarity(current_embedding, existing_embedding) - if similarity > threshold: - logger.debug(f"Found similar task memory in current batch with similarity: {similarity:.3f}") - return True - return False - - @staticmethod - def _calculate_cosine_similarity(embedding1: List[float], embedding2: List[float]) -> float: - """Calculate cosine similarity""" - try: - import numpy as np - - vec1 = np.array(embedding1) - vec2 = np.array(embedding2) - - # Calculate cosine similarity - dot_product = np.dot(vec1, vec2) - norm1 = np.linalg.norm(vec1) - norm2 = np.linalg.norm(vec2) - - if norm1 == 0 or norm2 == 0: - return 0.0 - - return dot_product / (norm1 * norm2) - - except Exception as e: - logger.error(f"Error calculating cosine similarity: {e}") - return 0.0 diff --git a/reme/extension/procedural_memory/summary/memory_validation.py b/reme/extension/procedural_memory/summary/memory_validation.py deleted file mode 100644 index a92b158f..00000000 --- a/reme/extension/procedural_memory/summary/memory_validation.py +++ /dev/null @@ -1,139 +0,0 @@ -"""Memory validation operation for task memory quality control. - -This module provides operations to validate the quality of extracted task -memories using LLM-based evaluation, ensuring only high-quality memories -are stored. -""" - -import json -import re -from typing import List, Dict, Any - -from loguru import logger - -from ....core.enumeration import Role -from ....core.op import BaseOp -from ....core.schema.message import Message -from ....core.schema.memory_node import MemoryNode - - -class MemoryValidation(BaseOp): - """Validate quality of extracted task memories. - - This operation uses LLM-based evaluation to assess the quality of extracted - task memories, filtering out low-quality or invalid memories based on - validation scores and criteria. - """ - - async def execute(self): - """Validate quality of extracted task memories""" - - task_memories: List[MemoryNode] = [] - task_memories.extend(self.context.get("success_task_memories", [])) - task_memories.extend(self.context.get("failure_task_memories", [])) - task_memories.extend(self.context.get("comparative_task_memories", [])) - - if not task_memories: - logger.info("No task memories found for validation") - return - - logger.info(f"Validating {len(task_memories)} extracted task memories") - - # Validate task memories - validated_task_memories = [] - - for task_memory in task_memories: - 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) - else: - reason = validation_result.get("reason", "Unknown reason") if validation_result else "Validation failed" - logger.warning(f"Task memory validation failed: {reason}") - - logger.info(f"Validated {len(validated_task_memories)} out of {len(task_memories)} task memories") - - # Update context - self.context.response.answer = json.dumps([x.model_dump() for x in validated_task_memories]) - self.context.response.metadata["memory_list"] = validated_task_memories - - async def _validate_single_task_memory(self, task_memory: MemoryNode) -> Dict[str, Any]: - """Validate single task memory""" - validation_info = await self._llm_validate_task_memory(task_memory) - logger.info(f"Validating: {validation_info}") - return validation_info - - async def _llm_validate_task_memory(self, task_memory: MemoryNode) -> Dict[str, Any]: - """Validate task memory using LLM""" - try: - prompt = self.prompt_format( - prompt_name="task_memory_validation_prompt", - condition=task_memory.when_to_use, - task_memory_content=task_memory.content, - ) - - def parse_validation(message: Message) -> Dict[str, Any]: - try: - response_content = message.content - - # Parse validation result - # Extract JSON blocks - json_pattern = r"```json\s*([\s\S]*?)\s*```" - json_blocks = re.findall(json_pattern, response_content) - - parsed: Dict[str, Any] = {} - if json_blocks: - raw_json = json_blocks[0] - try: - parsed = json.loads(raw_json) - except json.JSONDecodeError as json_err: - logger.warning( - f"JSONDecodeError in task_memory_validation, fallback to regex parse: {json_err}", - ) - is_valid_match = re.search(r'"is_valid"\s*:\s*(true|false)', raw_json, re.IGNORECASE) - score_match = re.search(r'"score"\s*:\s*([0-9]+(?:\.[0-9]+)?)', raw_json) - - if is_valid_match: - parsed["is_valid"] = is_valid_match.group(1).lower() == "true" - if score_match: - parsed["score"] = float(score_match.group(1)) - - is_valid = parsed.get("is_valid", True) - score = parsed.get("score", 0.5) - - # Set validation threshold - validation_threshold = self.context.get("validation_threshold", 0.5) - - return { - "is_valid": is_valid and score >= validation_threshold, - "score": score, - "feedback": response_content, - "reason": ( - "" - if (is_valid and score >= validation_threshold) - else f"Low validation score ({score:.2f}) or marked as invalid" - ), - } - - except Exception as e_inner: - logger.exception(f"Error parsing validation response: {e_inner}") - return { - "is_valid": False, - "score": 0.0, - "feedback": "", - "reason": f"Parse error: {str(e_inner)}", - } - - return await self.llm.chat( - messages=[Message(role=Role.USER, content=prompt)], - callback_fn=parse_validation, - ) - - except Exception as e: - logger.error(f"LLM validation failed: {e}") - return { - "is_valid": False, - "score": 0.0, - "feedback": "", - "reason": f"LLM validation error: {str(e)}", - } diff --git a/reme/extension/procedural_memory/summary/success_extraction.py b/reme/extension/procedural_memory/summary/success_extraction.py deleted file mode 100644 index 4e45459f..00000000 --- a/reme/extension/procedural_memory/summary/success_extraction.py +++ /dev/null @@ -1,95 +0,0 @@ -"""Success extraction operation for task memory generation. - -This module provides operations to extract task memories from successful -trajectories, identifying patterns and strategies that lead to success. -""" - -from typing import List - -from loguru import logger - -from ....core.enumeration import MemoryType, Role -from ....core.op import BaseOp -from ....core.schema.memory_node import MemoryNode -from ....core.schema.message import Message, Trajectory -from ..utils import ( - get_trajectory_context, - merge_messages_content, - parse_json_experience_response, -) - - -class SuccessExtraction(BaseOp): - """Extract task memories from successful trajectories. - - This operation analyzes successful trajectories (or their segments) to - extract reusable patterns, strategies, and best practices that can be - applied to similar future tasks. - """ - - async def execute(self): - """Extract task memories from successful trajectories""" - success_trajectories: List[Trajectory] = self.context.success_trajectories - - if not success_trajectories: - logger.info("No success trajectories found for extraction") - return - - logger.info(f"Extracting task memories from {len(success_trajectories)} successful trajectories") - - success_task_memories = [] - - # Process trajectories - for trajectory in success_trajectories: - if "segments" in trajectory.metadata: - # Process segmented step sequences - for segment in trajectory.metadata["segments"]: - task_memories = await self._extract_success_task_memory_from_steps(segment, trajectory) - success_task_memories.extend(task_memories) - else: - # Process entire 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") - - # Add task memories to context - self.context.success_task_memories = success_task_memories - - async def _extract_success_task_memory_from_steps( - self, - steps: List[Message], - trajectory: Trajectory, - ) -> List[MemoryNode]: - """Extract task memory from successful step sequences""" - step_content = merge_messages_content(steps) - context = get_trajectory_context(trajectory, steps) - - prompt = self.prompt_format( - prompt_name="success_step_task_memory_prompt", - query=trajectory.metadata.get("query", ""), - step_sequence=step_content, - context=context, - outcome="successful", - ) - - def parse_task_memories(message: Message) -> list[MemoryNode]: - task_memories_data = parse_json_experience_response(message.content) # extract content - task_memories = [] - - for tm_data in task_memories_data: - task_memory = MemoryNode( - memory_type=MemoryType.PROCEDURAL, - when_to_use=tm_data.get("when_to_use", tm_data.get("condition", "")), - content=tm_data.get("experience", ""), - author=getattr(self.llm, "model_name", "system"), - metadata=tm_data, - ) - task_memories.append(task_memory) - - return task_memories - - return await self.llm.chat( - messages=[Message(role=Role.USER, content=prompt)], - callback_fn=parse_task_memories, - ) diff --git a/reme/extension/procedural_memory/summary/trajectory_preprocess.py b/reme/extension/procedural_memory/summary/trajectory_preprocess.py deleted file mode 100644 index e099f483..00000000 --- a/reme/extension/procedural_memory/summary/trajectory_preprocess.py +++ /dev/null @@ -1,59 +0,0 @@ -"""Trajectory preprocessing operation for task memory generation. - -This module provides operations to preprocess and classify trajectories -into success and failure categories based on score thresholds. -""" - -from typing import Dict, List - -from loguru import logger - -from ....core.op import BaseOp -from ....core.schema.message import Trajectory - - -class TrajectoryPreprocess(BaseOp): - """Preprocess trajectories: validate and classify by success/failure. - - This operation classifies trajectories into success and failure categories - based on score thresholds, preparing them for downstream memory extraction - operations. - """ - - async def 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] - - # Classify trajectories - classified = self._classify_trajectories(trajectories) - logger.info( - f"Classified trajectories - Success: {len(classified['success'])}, " - f"Failure: {len(classified['failure'])}, All: {len(classified['all'])}", - ) - - # Set context for downstream operators - self.context.success_trajectories = classified["success"] - self.context.failure_trajectories = classified["failure"] - self.context.all_trajectories = classified["all"] - - def _classify_trajectories(self, trajectories: List[Trajectory]) -> Dict[str, List[Trajectory]]: - """Classify trajectories based on score threshold""" - success_trajectories = [] - failure_trajectories = [] - - success_threshold = self.context.get("success_threshold", 1.0) - - for traj in trajectories: - is_success = traj.score >= success_threshold - - if is_success: - success_trajectories.append(traj) - else: - failure_trajectories.append(traj) - - return { - "success": success_trajectories, - "failure": failure_trajectories, - "all": trajectories, - } diff --git a/reme/extension/procedural_memory/summary/trajectory_segmentation.py b/reme/extension/procedural_memory/summary/trajectory_segmentation.py deleted file mode 100644 index 9673ab76..00000000 --- a/reme/extension/procedural_memory/summary/trajectory_segmentation.py +++ /dev/null @@ -1,139 +0,0 @@ -"""Trajectory segmentation operation for task memory generation. - -This module provides operations to segment trajectories into meaningful step -sequences that can be used for more granular memory extraction. -""" - -import json -import re -from typing import List - -from loguru import logger - -from ....core.enumeration import Role -from ....core.op import BaseOp -from ....core.schema.message import Message, Trajectory - - -class TrajectorySegmentation(BaseOp): - """Segment trajectories into meaningful step sequences. - - This operation uses LLM to identify natural breakpoints in trajectories, - allowing for more granular analysis and memory extraction from specific - segments rather than entire trajectories. - """ - - async def execute(self): - """Segment trajectories into meaningful steps""" - # Get trajectories from context - all_trajectories: List[Trajectory] = self.context.get("all_trajectories", []) - success_trajectories: List[Trajectory] = self.context.get("success_trajectories", []) - failure_trajectories: List[Trajectory] = self.context.get("failure_trajectories", []) - - if not all_trajectories: - logger.warning("No trajectories found in context") - return - - # Determine which trajectories to segment - target_trajectories = self._get_target_trajectories( - all_trajectories, - success_trajectories, - failure_trajectories, - ) - - # Add segmentation info to trajectories - segmented_count = 0 - for trajectory in target_trajectories: - segments = await self._llm_segment_trajectory(trajectory) - trajectory.metadata["segments"] = segments - segmented_count += 1 - - logger.info(f"Segmented {segmented_count} trajectories") - - # Update context with segmented trajectories - - def _get_target_trajectories( - self, - all_trajectories: List[Trajectory], - success_trajectories: List[Trajectory], - failure_trajectories: List[Trajectory], - ) -> List[Trajectory]: - """Determine which trajectories to segment based on configuration""" - segment_target = self.context.get("segment_target", "all") - - if segment_target == "success": - return success_trajectories - elif segment_target == "failure": - return failure_trajectories - else: - return all_trajectories - - async def _llm_segment_trajectory(self, trajectory: Trajectory) -> List[List[Message]]: - """Use LLM for trajectory segmentation""" - trajectory_content = self._format_trajectory_content(trajectory) - - prompt = self.prompt_format( - prompt_name="step_segmentation_prompt", - query=trajectory.metadata.get("query", ""), - trajectory_content=trajectory_content, - total_steps=len(trajectory.messages), - ) - - def parse_segmentation(message: Message) -> List[List[Message]]: - content = message.content - segment_points = self._parse_segmentation_response(content) - - # Segment trajectory based on segmentation points - segments = [] - start_idx = 0 - - for end_idx in segment_points: - if start_idx < end_idx <= len(trajectory.messages): - segments.append(trajectory.messages[start_idx:end_idx]) - start_idx = end_idx - - # Add remaining steps - if start_idx < len(trajectory.messages): - segments.append(trajectory.messages[start_idx:]) - - return segments if segments else [trajectory.messages] - - return await self.llm.chat( - messages=[Message(role=Role.USER, content=prompt)], - callback_fn=parse_segmentation, - default_value=[trajectory.messages], - ) - - @staticmethod - def _format_trajectory_content(trajectory: Trajectory) -> str: - """Format trajectory content for LLM processing""" - content = "" - for i, step in enumerate(trajectory.messages): - content += f"Step {i + 1} ({step.role.value}):\n{step.content}\n\n" - return content - - @staticmethod - def _parse_segmentation_response(response: str) -> List[int]: - """Parse segmentation response from LLM""" - segment_points = [] - - # Try to extract JSON format - json_pattern = r"```json\s*([\s\S]*?)\s*```" - json_blocks = re.findall(json_pattern, response) - - if json_blocks: - try: - parsed = json.loads(json_blocks[0]) - if isinstance(parsed, dict) and "segment_points" in parsed: - segment_points = parsed["segment_points"] - elif isinstance(parsed, list): - segment_points = parsed - except json.JSONDecodeError: - pass - - # Fallback: extract numbers - if not segment_points: - numbers = re.findall(r"\b\d+\b", response) - segment_points = [int(num) for num in numbers if int(num) > 0] - - return sorted(list(set(segment_points))) diff --git a/reme/extension/procedural_memory/utils.py b/reme/extension/procedural_memory/utils.py deleted file mode 100644 index e6130bb1..00000000 --- a/reme/extension/procedural_memory/utils.py +++ /dev/null @@ -1,142 +0,0 @@ -"""Utility functions for processing and formatting LLM-related message data.""" - -import json -import re -from loguru import logger - -from ...core.enumeration import Role -from ...core.schema.message import Message, Trajectory - - -def merge_messages_content(messages: list[Message | dict]) -> str: - """Merge messages content into a formatted string representation. - - This function processes a list of messages (either Message objects or dicts) - and formats them into a structured string. Different message roles are - formatted differently: - - ASSISTANT: Includes reasoning content, main content, and tool calls - - USER: Includes the user content - - TOOL: Includes tool call results - - Each message is prefixed with a step number (starting from 0) to indicate - its position in the conversation sequence. - - Args: - messages: List of Message objects or dictionaries to merge. If a dict - is provided, it will be converted to a Message object. - - Returns: - Formatted string representation of all messages with step numbers. - Each message is separated by newlines and includes role information. - - Example: - ```python - messages = [ - Message(role=Role.USER, content="What's the weather?"), - Message(role=Role.ASSISTANT, content="Let me check", - tool_calls=[ToolCall(name="get_weather", arguments={})]) - ] - result = merge_messages_content(messages) - # Returns formatted string with step numbers and role information - ``` - """ - content_collector = [] - for i, message in enumerate(messages): - if isinstance(message, dict): - message = Message(**message) - - if message.role is Role.ASSISTANT: - line = ( - f"### step.{i} role={message.role.value} content=\n{message.reasoning_content}\n\n{message.content}\n" - ) - if message.tool_calls: - for tool_call in message.tool_calls: - line += f" - tool call={tool_call.name}\n params={tool_call.arguments}\n" - content_collector.append(line) - - elif message.role is Role.USER: - line = f"### step.{i} role={message.role.value} content=\n{message.content}\n" - content_collector.append(line) - - elif message.role is Role.TOOL: - line = f"### step.{i} role={message.role.value} tool call result=\n{message.content}\n" - content_collector.append(line) - - return "\n".join(content_collector) - - -def parse_json_experience_response(response: str) -> list[dict]: - """Parse JSON formatted experience response""" - try: - # Extract JSON blocks - json_pattern = r"```json\s*([\s\S]*?)\s*```" - json_blocks = re.findall(json_pattern, response) - - if json_blocks: - parsed = json.loads(json_blocks[0]) - - # Handle array format - if isinstance(parsed, list): - experiences = [] - for exp_data in parsed: - if isinstance(exp_data, dict) and ( - ("when_to_use" in exp_data and "experience" in exp_data) - or ("condition" in exp_data and "experience" in exp_data) - ): - experiences.append(exp_data) - - return experiences - - # Handle single object - elif isinstance(parsed, dict) and ( - ("when_to_use" in parsed and "experience" in parsed) - or ("condition" in parsed and "experience" in parsed) - ): - return [parsed] - - # Fallback: try to parse entire response - parsed = json.loads(response) - if isinstance(parsed, list): - return parsed - elif isinstance(parsed, dict): - return [parsed] - - except json.JSONDecodeError as e: - logger.warning(f"Failed to parse JSON experience response: {e}") - - return [] - - -def get_trajectory_context(trajectory: Trajectory, step_sequence: list[Message]) -> str: - """Get context of step sequence within trajectory""" - try: - # Find position of step sequence in trajectory - start_idx = 0 - for i, step in enumerate(trajectory.messages): - if step == step_sequence[0]: - start_idx = i - break - - # Extract before and after context - context_before = trajectory.messages[max(0, start_idx - 2) : start_idx] - context_after = trajectory.messages[start_idx + len(step_sequence) : start_idx + len(step_sequence) + 2] - - context = f"Query: {trajectory.metadata.get('query', 'N/A')}\n" - - if context_before: - context += ( - "Previous steps:\n" - + "\n".join( - [f"- {step.content[:100]}..." for step in context_before], - ) - + "\n" - ) - - if context_after: - context += "Following steps:\n" + "\n".join([f"- {step.content[:100]}..." for step in context_after]) - - return context - - except Exception as e: - logger.error(f"Error getting trajectory context: {e}") - return f"Query: {trajectory.metadata.get('query', 'N/A')}" diff --git a/reme/extension/simple_chat.py b/reme/extension/simple_chat.py deleted file mode 100644 index 04a36656..00000000 --- a/reme/extension/simple_chat.py +++ /dev/null @@ -1,60 +0,0 @@ -"""Simple chat for test.""" - -from loguru import logger - -from ..core.enumeration import Role -from ..core.op import BaseTool -from ..core.schema import Message, ToolCall - - -class SimpleChat(BaseTool): - """Simple chat agent that handles non-streaming conversations.""" - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "simple chat agent", - "parameters": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "query", - }, - "messages": { - "type": "array", - "items": { - "type": "object", - "properties": { - "role": { - "type": "string", - "description": "role", - }, - "content": { - "type": "string", - "description": "content", - }, - }, - "required": ["role", "content"], - }, - }, - }, - "required": [], - }, - }, - ) - - async def execute(self): - if "query" in self.context: - messages = [ - Message(role=Role.SYSTEM, content="You are a helpful assistant."), - Message(role=Role.USER, content=self.context.query), - ] - elif "messages" in self.context: - messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages if m] - else: - raise ValueError("query or messages must be provided!") - logger.info(f"messages={messages}") - assistant_message = await self.llm.chat(messages=messages) - logger.info(f"assistant_message={assistant_message.simple_dump()}") - return assistant_message.content diff --git a/reme/extension/stream_chat.py b/reme/extension/stream_chat.py deleted file mode 100644 index dd0aeda7..00000000 --- a/reme/extension/stream_chat.py +++ /dev/null @@ -1,62 +0,0 @@ -"""Streaming chat for test.""" - -from loguru import logger - -from ..core.enumeration import Role, ChunkEnum -from ..core.op import BaseTool -from ..core.schema import Message, ToolCall - - -class StreamChat(BaseTool): - """Streaming chat agent that handles real-time conversation streaming.""" - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "simple chat agent", - "parameters": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "query", - }, - "messages": { - "type": "array", - "items": { - "type": "object", - "properties": { - "role": { - "type": "string", - "description": "role", - }, - "content": { - "type": "string", - "description": "content", - }, - }, - "required": ["role", "content"], - }, - }, - }, - "required": [], - }, - }, - ) - - async def execute(self): - """Execute streaming chat operation with query or messages.""" - if "query" in self.context: - messages = [ - Message(role=Role.SYSTEM, content="You are a helpful assistant."), - Message(role=Role.USER, content=self.context.query), - ] - elif "messages" in self.context: - messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages if m] - else: - raise ValueError("query or messages must be provided!") - logger.info(f"messages={messages}") - - async for stream_chunk in self.llm.stream_chat(messages): - if stream_chunk.chunk_type in [ChunkEnum.ANSWER, ChunkEnum.THINK, ChunkEnum.ERROR, ChunkEnum.TOOL]: - await self.context.add_stream_chunk(stream_chunk) diff --git a/reme/extension/test_op.py b/reme/extension/test_op.py deleted file mode 100644 index acc1ddcf..00000000 --- a/reme/extension/test_op.py +++ /dev/null @@ -1,15 +0,0 @@ -"""Test workflow operations.""" - -from loguru import logger - -from ..core.op import BaseOp - - -class TestOp(BaseOp): - """Test operation for workflow testing.""" - - async def execute(self): - logger.info("delete start") - # await self.vector_store.delete_all() - await self.vector_store.delete("123") - logger.info("delete end") diff --git a/reme/extension/translate_ts.py b/reme/extension/translate_ts.py deleted file mode 100644 index eb032386..00000000 --- a/reme/extension/translate_ts.py +++ /dev/null @@ -1,55 +0,0 @@ -"""Translate operation for translating text.""" - -from pathlib import Path - -from loguru import logger - -from ..core.enumeration import Role -from ..core.op import BaseOp -from ..core.schema import Message - - -class TranslateTs(BaseOp): - """ - Translate operation for translating text. - - reme2 backend=cmd cmd.flow="TranslateTs()" cmd.params.target_dir="" - """ - - async def execute_single_file(self, ts_file: Path): - """Translate a single ts file.""" - ts_code = ts_file.read_text(encoding="utf-8") - logger.info(f"Translating {ts_file}") - - def parse_python(assistant_message: Message): - python_code = assistant_message.content - assert "```python" in python_code, "Invalid python code" - python_code = python_code.split("```python", 1)[1] - python_code_split = python_code.split("```") - python_code = "```".join(python_code_split[:-1]) - return python_code.strip(), assistant_message.content.strip() - - output = await self.llm.chat( - messages=[ - Message( - role=Role.USER, - content=self.prompt_format(prompt_name="translate_prompt", ts_code=ts_code), - ), - ], - callback_fn=parse_python, - ) - - ts_file.with_suffix(".py").write_text(output[0], encoding="utf-8") - ts_file.with_suffix(".txt").write_text(output[1], encoding="utf-8") - logger.info(f"Translate {ts_file} complete.") - - async def execute(self): - """Execute the operation.""" - target_dir = Path(self.context.target_dir) - ts_files = [p for p in target_dir.rglob("*.ts") if p.is_file() and not p.name.endswith(".test.ts")] - logger.info(f"Translating {target_dir}, finding {len(ts_files)} ts files") - - for ts_file in ts_files: - self.submit_async_task(self.execute_single_file, ts_file) - - await self.join_async_tasks() diff --git a/reme/memory/__init__.py b/reme/memory/__init__.py deleted file mode 100644 index e209a3fe..00000000 --- a/reme/memory/__init__.py +++ /dev/null @@ -1,11 +0,0 @@ -"""memory""" - -from . import file_based -from . import vector_tools -from . import vector_based - -__all__ = [ - "file_based", - "vector_tools", - "vector_based", -] diff --git a/reme/memory/file_based/__init__.py b/reme/memory/file_based/__init__.py deleted file mode 100644 index e4c392e5..00000000 --- a/reme/memory/file_based/__init__.py +++ /dev/null @@ -1,13 +0,0 @@ -"""File-based Memory Module.""" - -from . import components -from . import tools -from . import utils -from .reme_in_memory_memory import ReMeInMemoryMemory - -__all__ = [ - "tools", - "utils", - "components", - "ReMeInMemoryMemory", -] diff --git a/reme/memory/file_based/components/__init__.py b/reme/memory/file_based/components/__init__.py deleted file mode 100644 index 86574790..00000000 --- a/reme/memory/file_based/components/__init__.py +++ /dev/null @@ -1,15 +0,0 @@ -"""components""" - -from .compactor import Compactor -from .context_checker import ContextChecker -from .summarizer import Summarizer -from .tool_result_compactor import ToolResultCompactor -from .cli import CliAgent - -__all__ = [ - "Compactor", - "Summarizer", - "ContextChecker", - "ToolResultCompactor", - "CliAgent", -] diff --git a/reme/memory/file_based/components/cli.py b/reme/memory/file_based/components/cli.py deleted file mode 100644 index 138a3aad..00000000 --- a/reme/memory/file_based/components/cli.py +++ /dev/null @@ -1,326 +0,0 @@ -"""CLI component for interactive chat using agentscope-based memory tools.""" - -import asyncio -from datetime import datetime -from pathlib import Path -import zoneinfo - -from agentscope.agent import ReActAgent -from agentscope.message import Msg, TextBlock -from agentscope.pipeline import stream_printing_messages -from agentscope.tool import Toolkit, ToolResponse - -from .compactor import Compactor -from .context_checker import ContextChecker -from .summarizer import Summarizer -from ..tools import FileIO, MemorySearch -from ....core.op import BaseOp -from ....core.utils import format_messages -from ....core.utils import get_logger - -logger = get_logger() -# name + desc + "{working_dir}/skills/{skill_name}/SKILL.md" - -_DEFAULT_AGENT_SKILL_INSTRUCTION = ( - "# Agent Skills\n" - "The agent skills are a collection of folds of instructions, scripts, " - "and resources that you can load dynamically to improve performance " - "on specialized tasks. Each agent skill has a `SKILL.md` file in its " - "folder that describes how to use the skill. If you want to use a " - "skill, you MUST read its `SKILL.md` file carefully." -) - -_DEFAULT_AGENT_SKILL_TEMPLATE = """## {name} -{description} -Check "{dir}/SKILL.md" for how to use this skill""" - - -class CliAgent(BaseOp): - """CLI agent for interactive chat with memory management.""" - - def __init__( - self, - working_dir: str, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - context_window_tokens: int = 128000, - reserve_tokens: int = 36000, - keep_recent_tokens: int = 20000, - language: str = "zh", - timezone: str | None = None, - **kwargs, - ): - super().__init__(**kwargs) - self.working_dir: str = working_dir - Path(self.working_dir).mkdir(parents=True, exist_ok=True) - self.vector_weight: float = vector_weight - self.candidate_multiplier: float = candidate_multiplier - self.context_window_tokens: int = context_window_tokens - self.reserve_tokens: int = reserve_tokens - self.keep_recent_tokens: int = keep_recent_tokens - self.language: str = language - self.timezone: str | None = timezone - - # Initialize message history - self.messages: list[Msg] = [] - self.previous_summary: str = "" - self.summary_tasks: list[asyncio.Task] = [] - - def add_summary_task(self, messages: list[Msg]): - """Add summary task to queue.""" - remaining_tasks = [] - for task in self.summary_tasks: - if task.done(): - exc = task.exception() - if exc is not None: - logger.exception(f"Summary task failed: {exc}") - else: - result = task.result() - logger.info(f"Summary task completed: {result}") - else: - remaining_tasks.append(task) - self.summary_tasks = remaining_tasks - - # Create a toolkit for the summarizer - toolkit = self._create_file_toolkit() - - # Create summarizer instance - memory_path = Path(self.working_dir) / "memory" - summarizer = Summarizer( - working_dir=self.working_dir, - memory_dir=str(memory_path), - memory_compact_threshold=int(self.context_window_tokens * 0.7), - token_counter=self.as_token_counter, - toolkit=toolkit, - as_llm=self.as_llm, - as_llm_formatter=self.as_llm_formatter, - language=self.language if self.language == "zh" else "", - console_enabled=False, # We disable the terminal printing to avoid messy outputs - timezone=self.timezone, - ) - - # Create summary task - summary_task = asyncio.create_task( - summarizer.call( - messages=messages, - service_context=self.service_context, - ), - ) - self.summary_tasks.append(summary_task) - - def _create_file_toolkit(self): - """Create a toolkit with file operations.""" - - toolkit = Toolkit() - file_io = FileIO(working_dir=self.working_dir) - toolkit.register_tool_function(file_io.read_file) - toolkit.register_tool_function(file_io.write_file) - toolkit.register_tool_function(file_io.edit_file) - - return toolkit - - async def new(self) -> str: - """Reset conversation history using summary.""" - if not self.messages: - self.messages.clear() - self.previous_summary = "" - return "No history to reset." - - self.add_summary_task(self.messages) - - self.messages.clear() - self.previous_summary = "" - return "History saved to memory files and reset." - - async def context_check(self) -> dict: - """Check if messages exceed token limits.""" - # Create context checker - checker = ContextChecker( - memory_compact_threshold=self.context_window_tokens - self.reserve_tokens, - memory_compact_reserve=self.keep_recent_tokens, - token_counter=self.as_token_counter, - ) - - return await checker.call( - messages=self.messages, - service_context=self.service_context, - ) - - async def compact(self, force_compact: bool = False) -> str: - """Compact history then reset.""" - if not self.messages: - return "No history to compact." - - # Check and find cut point - messages_to_compact, messages_to_keep, _ = await self.context_check() - tokens_before = len(self.messages) - - if force_compact: - messages_to_summarize = self.messages - left_messages = [] - elif not messages_to_compact: - return "History is within token limits, no compaction needed." - else: - messages_to_summarize = messages_to_compact - left_messages = messages_to_keep - - # Create compactor - compactor = Compactor( - memory_compact_threshold=self.context_window_tokens - self.reserve_tokens, - token_counter=self.as_token_counter, - as_llm=self.as_llm, - as_llm_formatter=self.as_llm_formatter, - language=self.language if self.language == "zh" else "", - console_enabled=False, # We disable the terminal printing to avoid messy outputs - timezone=self.timezone, - ) - - summary_content = await compactor.call( - messages=messages_to_summarize, - previous_summary=self.previous_summary, - service_context=self.service_context, - ) - - self.add_summary_task(messages=messages_to_summarize) - - # Assemble final messages - self.messages = left_messages - self.previous_summary = summary_content - - return f"History compacted from {tokens_before} messages." - - def format_history(self) -> str: - """Format history messages.""" - return format_messages( - messages=self.messages, - add_index=False, - add_reasoning=False, - strip_markdown_headers=False, - ) - - async def _build_messages(self, query: str) -> list[Msg]: - """Build system prompt message.""" - tz = zoneinfo.ZoneInfo(self.timezone) if self.timezone else None - current_time = datetime.now(tz).strftime("%Y-%m-%d %H:%M:%S %A") - - # Create system prompt - system_prompt = self.prompt_format( - "system_prompt", - workspace_dir=self.working_dir, - current_time=current_time, - has_previous_summary=bool(self.previous_summary), - previous_summary=self.previous_summary or "", - ) - - logger.info(f"[{self.__class__.__name__}] system_prompt: {system_prompt}") - - # Build message list - messages = [Msg(name="system", role="system", content=system_prompt)] - messages.extend(self.messages) - messages.append(Msg(name="user", role="user", content=query)) - - return messages - - async def memory_search(self, query: str, max_results: int = 5, min_score: float = 0.1) -> ToolResponse: - """ - Mandatory recall step: semantically search MEMORY.md + memory/*.md (and optional session transcripts) - before answering questions about prior work, decisions, dates, people, preferences, or todos; - returns top snippets with path + lines. - - Args: - query: The semantic search query to find relevant memory snippets - max_results: Maximum number of search results to return (optional), default is 5 - min_score: Minimum similarity score threshold for results (optional), default is 0.1 - - Returns: - Search results as formatted string - """ - search_tool = MemorySearch( - vector_weight=self.vector_weight, - candidate_multiplier=self.candidate_multiplier, - ) - search_result = await search_tool.call( - query=query, - max_results=max_results, - min_score=min_score, - service_context=self.service_context, - ) - return ToolResponse( - content=[ - TextBlock( - type="text", - text=search_result, - ), - ], - ) - - async def execute(self): - """Execute the agent.""" - _ = await self.compact(force_compact=False) - - # Build messages for the agent - query = self.context.query - messages = await self._build_messages(query) - - toolkit = self._create_file_toolkit() - # Register memory search tool - toolkit.register_tool_function(self.memory_search) - - # Create the ReAct agent - agent = ReActAgent( - name="reme_cli_agent", - model=self.as_llm, - sys_prompt=messages[0].content, # System prompt - formatter=self.as_llm_formatter, - toolkit=toolkit, - ) - - # We disable the terminal printing to avoid messy outputs - agent.set_console_output_enabled(False) - - self.messages = messages[1:] # remove the first SYSTEM message - agent.memory.content.clear() - - # Stream processing state - in_thinking = False - in_answer = False - - # obtain the printing messages from the agent in a streaming way - last_text_content = "" - last_think_content = "" - async for msg, last in stream_printing_messages( - agents=[agent], - coroutine_task=agent(self.messages), - ): - # print(msg, last) - content_blocks = msg.get_content_blocks() - for block in content_blocks: - if block["type"] == "thinking": - if not in_thinking and len(block["thinking"]) > len(last_think_content): - print("\033[90m\nThinking: ", end="", flush=True) - in_thinking = True - print(block["thinking"][len(last_think_content) :], end="", flush=True) - last_think_content = block["thinking"] - elif block["type"] == "text": - if in_thinking: - print("\033[0m") # reset color after thinking - in_thinking = False - if not in_answer: - print("\nRemy: ", end="", flush=True) - in_answer = True - print(block["text"][len(last_text_content) :], end="", flush=True) - last_text_content = block["text"] - elif block["type"] == "tool_use": - if in_thinking: - print("\033[0m") # reset color after thinking - in_thinking = False - if last: - print(f"\033[36m -> Executing Tool: name={block['name']}, input={block['input']}\033[0m") - elif block["type"] == "tool_result": - if last: - last_think_content = "" # reset for further thinking - print(f"\033[36m -> Tool Result for `{block['name']}`: {block['output'][0]['text']}\033[0m") - else: - print(f"Unknown block type: {block['type']}") - if last: - self.messages.append(msg) diff --git a/reme/memory/file_based/components/compactor.py b/reme/memory/file_based/components/compactor.py deleted file mode 100644 index e8fe9383..00000000 --- a/reme/memory/file_based/components/compactor.py +++ /dev/null @@ -1,115 +0,0 @@ -"""Compactor module for memory compaction operations.""" - -from agentscope.agent import ReActAgent -from agentscope.message import Msg - -from ..utils import AsMsgHandler -from ....core.op import BaseOp -from ....core.utils import get_logger - -logger = get_logger() - - -def _is_valid_summary(content: str) -> bool: - """Check if the summary content is valid. - - Args: - content: The summary content to validate. - - Returns: - True if valid, False otherwise. - """ - if not content or not content.strip(): - return False - if "##" not in content: - return False - return True - - -class Compactor(BaseOp): - """Compactor class for compacting memory messages.""" - - def __init__( - self, - memory_compact_threshold: int, - console_enabled: bool = False, - return_dict: bool = False, - add_thinking_block: bool = True, - extra_instruction: str = "", - **kwargs, - ): - super().__init__(**kwargs) - self.memory_compact_threshold: int = memory_compact_threshold - self.console_enabled: bool = console_enabled - self.return_dict: bool = return_dict - self.add_thinking_block: bool = add_thinking_block - self.extra_instruction: str = extra_instruction - - # pylint: disable=too-many-return-statements - async def execute(self): - messages: list[Msg] = self.context.get("messages", []) - previous_summary: str = self.context.get("previous_summary", "") - - if not messages: - if self.return_dict: - return {"user_message": "", "history_compact": "", "is_valid": False} - return "" - - msg_handler = AsMsgHandler(self.as_token_counter) - before_token_count = await msg_handler.count_msgs_token(messages) - history_formatted_str: str = await msg_handler.format_msgs_to_str( - messages=messages, - memory_compact_threshold=self.memory_compact_threshold, - include_thinking=self.add_thinking_block, - ) - after_token_count = await msg_handler.count_str_token(history_formatted_str) - logger.info(f"Compactor before_token_count={before_token_count} after_token_count={after_token_count}") - - if not history_formatted_str: - logger.warning(f"No history to compact. messages={messages}") - if self.return_dict: - return {"user_message": "", "history_compact": "", "is_valid": False} - return "" - - agent = ReActAgent( - name="reme_compactor", - model=self.as_llm, - sys_prompt=self.get_prompt("system_prompt"), - formatter=self.as_llm_formatter, - ) - agent.set_console_output_enabled(self.console_enabled) - - if previous_summary: - user_message: str = ( - f"# conversation\n{history_formatted_str}\n\n" - f"# previous-summary\n{previous_summary}\n\n" + self.get_prompt("update_user_message") - ) - else: - user_message: str = f"# conversation\n{history_formatted_str}\n\n" + self.get_prompt("initial_user_message") - - if self.extra_instruction: - user_message += f"\n\n# extra-instruction\n{self.extra_instruction}" - logger.info(f"Compactor sys_prompt={agent.sys_prompt} user_message={user_message}") - - compact_msg: Msg = await agent.reply( - Msg( - name="reme", - role="user", - content=user_message, - ), - ) - - history_compact: str = compact_msg.get_text_content() - is_valid: bool = _is_valid_summary(history_compact) - - if not is_valid: - logger.warning(f"Invalid summary result: {history_compact[:200]}...") - if self.return_dict: - return {"user_message": user_message, "history_compact": history_compact, "is_valid": False} - return "" - - logger.info(f"Compactor Result:\n{history_compact}") - - if self.return_dict: - return {"user_message": user_message, "history_compact": history_compact, "is_valid": True} - return history_compact diff --git a/reme/memory/file_based/components/context_checker.py b/reme/memory/file_based/components/context_checker.py deleted file mode 100644 index f5a49b2b..00000000 --- a/reme/memory/file_based/components/context_checker.py +++ /dev/null @@ -1,92 +0,0 @@ -"""ContextChecker module for checking context size and splitting messages.""" - -from agentscope.message import Msg - -from ..utils import AsMsgHandler -from ....core.op import BaseOp -from ....core.utils import get_logger - -logger = get_logger() - - -class ContextChecker(BaseOp): - """ - ContextChecker class for checking context size and splitting messages. - - This class analyzes conversation messages to determine if the context - exceeds the specified token threshold and splits messages into two groups: - those that should be compacted and those to keep in context. - - Attributes: - memory_compact_threshold (int): Token count threshold for triggering compaction. - memory_compact_reserve (int): Token count to reserve for recent messages. - """ - - def __init__( - self, - memory_compact_threshold: int, - memory_compact_reserve: int = 10000, - **kwargs, - ): - """ - Initialize the ContextChecker. - - Args: - memory_compact_threshold (int): Token count threshold for triggering - compaction. Messages exceeding this threshold will be split. - memory_compact_reserve (int): Token count to reserve for recent messages - to keep in context. Defaults to 10000 tokens. - **kwargs: Additional keyword arguments passed to BaseOp. - """ - super().__init__(**kwargs) - self.memory_compact_threshold: int = memory_compact_threshold - self.memory_compact_reserve: int = memory_compact_reserve - assert self.memory_compact_threshold > self.memory_compact_reserve - - async def execute(self) -> tuple[list[Msg], list[Msg], bool]: - """ - Execute context check and split messages. - - Retrieves messages from context and checks if they exceed the token - threshold. If so, splits them into messages to compact and messages - to keep. - - Context Parameters: - messages (list[Msg]): List of conversation messages to check. - Retrieved from self.context.get("messages", []). - - Returns: - tuple[list[Msg], list[Msg], bool]: A tuple containing: - - messages_to_compact (list[Msg]): Older messages that should - be compacted/summarized. - - messages_to_keep (list[Msg]): Recent messages to keep in context. - - is_valid (bool): True if the split is valid (tool calls aligned), - False if splitting would break conversation integrity. - - Note: - - Returns ([], messages, True) if no compaction is needed. - - Ensures conversation pairs (user-assistant) are not split. - - is_valid=False indicates tool_use and tool_result are misaligned. - """ - messages: list[Msg] = self.context.get("messages", []) - - if not messages: - logger.info("ContextChecker: No messages to check.") - return [], [], True - - msg_handler = AsMsgHandler(self.as_token_counter) - messages_to_compact, messages_to_keep, is_valid = await msg_handler.context_check( - messages=messages, - memory_compact_threshold=self.memory_compact_threshold, - memory_compact_reserve=self.memory_compact_reserve, - ) - - if messages_to_compact: - logger.info( - f"ContextChecker Result: " - f"to_compact={len(messages_to_compact)}, " - f"to_keep={len(messages_to_keep)}, " - f"is_valid={is_valid}", - ) - - return messages_to_compact, messages_to_keep, is_valid diff --git a/reme/memory/file_based/components/summarizer.py b/reme/memory/file_based/components/summarizer.py deleted file mode 100644 index 9733eaa5..00000000 --- a/reme/memory/file_based/components/summarizer.py +++ /dev/null @@ -1,97 +0,0 @@ -"""Summarizer module for memory summarization operations.""" - -import datetime -import zoneinfo - -from agentscope.agent import ReActAgent -from agentscope.message import Msg -from agentscope.tool import Toolkit - -from ..utils import AsMsgHandler -from ....core.op import BaseOp -from ....core.utils import get_logger - -logger = get_logger() - - -class Summarizer(BaseOp): - """Summarizer class for summarizing memory messages.""" - - def __init__( - self, - working_dir: str, - memory_dir: str, - memory_compact_threshold: int, - toolkit: Toolkit | None = None, - console_enabled: bool = False, - timezone: str | None = None, - add_thinking_block: bool = True, - **kwargs, - ): - super().__init__(**kwargs) - self.working_dir: str = working_dir - self.memory_dir: str = memory_dir - self.memory_compact_threshold: int = memory_compact_threshold - self.toolkit: Toolkit | None = toolkit - self.console_enabled: bool = console_enabled - self.timezone: str | None = timezone - self.add_thinking_block: bool = add_thinking_block - - def _get_current_datetime(self) -> datetime.datetime: - """Get current datetime with timezone, fallback to local time if timezone is invalid.""" - if self.timezone: - try: - return datetime.datetime.now(zoneinfo.ZoneInfo(self.timezone)) - except Exception as e: - logger.error(f"Invalid timezone: {self.timezone}, falling back to local time error={e}") - return datetime.datetime.now() - - async def execute(self): - messages: list[Msg] = self.context.get("messages", []) - - if not messages: - return "" - - msg_handler = AsMsgHandler(self.as_token_counter) - before_token_count = await msg_handler.count_msgs_token(messages) - history_formatted_str: str = await msg_handler.format_msgs_to_str( - messages=messages, - memory_compact_threshold=self.memory_compact_threshold, - include_thinking=self.add_thinking_block, - ) - after_token_count = await msg_handler.count_str_token(history_formatted_str) - logger.info(f"Summarizer before_token_count={before_token_count} after_token_count={after_token_count}") - - if not history_formatted_str: - logger.warning(f"No history to summarize. messages={messages}") - return "" - - agent = ReActAgent( - name="reme_summarizer", - model=self.as_llm, - sys_prompt="You are a helpful assistant.", - formatter=self.as_llm_formatter, - toolkit=self.toolkit, - ) - agent.set_console_output_enabled(self.console_enabled) - - user_message: str = f"# conversation\n{history_formatted_str}\n\n" + self.prompt_format( - "user_message", - date=self._get_current_datetime().strftime("%Y-%m-%d"), - working_dir=self.working_dir, - memory_dir=self.memory_dir, - ) - - summary_msg: Msg = await agent.reply( - Msg( - name="reme", - role="user", - content=user_message, - ), - ) - for i, (msg, _) in enumerate(agent.memory.content): - logger.info(f"Summarizer memory[{i}]: {msg.content}") - - history_summary: str = summary_msg.get_text_content() - logger.info(f"Summarizer Result:\n{history_summary}") - return history_summary diff --git a/reme/memory/file_based/components/tool_result_compactor.py b/reme/memory/file_based/components/tool_result_compactor.py deleted file mode 100644 index d80b2f61..00000000 --- a/reme/memory/file_based/components/tool_result_compactor.py +++ /dev/null @@ -1,161 +0,0 @@ -"""Tool Result Compactor: truncate large tool results and save full content to files.""" - -import os -import sys -import uuid -from datetime import datetime, timedelta -from pathlib import Path - -from agentscope.message import Msg - -from ..utils import truncate_text_output, DEFAULT_MAX_BYTES, TRUNCATION_NOTICE_MARKER -from ....core.op import BaseOp -from ....core.utils import get_logger - -logger = get_logger() - - -class ToolResultCompactor(BaseOp): - """Truncate large tool_result outputs and save full content to files.""" - - def __init__( - self, - tool_result_dir: str | Path, - retention_days: int = 3, - old_max_bytes: int = 3000, - recent_max_bytes: int = DEFAULT_MAX_BYTES, - recent_n: int = 1, - encoding: str = "utf-8", - **kwargs, - ): - super().__init__(**kwargs) - self.tool_result_dir = Path(tool_result_dir) - self.retention_days = retention_days - self.old_max_bytes = old_max_bytes - self.recent_max_bytes = recent_max_bytes - self.recent_n = recent_n - self.encoding = encoding - self.tool_result_dir.mkdir(parents=True, exist_ok=True) - - def _truncate(self, content: str, max_bytes: int) -> str: - if not content: - return content - - try: - if TRUNCATION_NOTICE_MARKER in content: - return truncate_text_output(content, max_bytes=max_bytes, encoding=self.encoding) - - if len(content.encode(self.encoding)) <= max_bytes + 100: - return content - - saved_path: str | None = None - fp = self.tool_result_dir / f"{uuid.uuid4().hex}.txt" - fp.write_text(content, encoding=self.encoding) - saved_path = str(fp) - - return truncate_text_output( - content, - 1, - content.count("\n") + 1, - max_bytes, - file_path=saved_path, - encoding=self.encoding, - ) - except Exception as e: - logger.warning("Failed to truncate content, returning original: %s", e) - return content - - def _compact(self, output: str | list[dict], max_bytes: int) -> str | list[dict]: - """Truncate output to max_bytes, saving full content to file if needed.""" - - if isinstance(output, str): - return self._truncate(output, max_bytes) - if isinstance(output, list): - for b in output: - if isinstance(b, dict) and b.get("type") == "text": - b["text"] = self._truncate(b.get("text", ""), max_bytes) - return output - - async def execute(self) -> list[Msg]: - """Process all messages, truncating large tool results.""" - messages: list[Msg] = self.context.get("messages", []) - if not messages: - return messages - - recent_n = 0 - for msg in reversed(messages): - if not isinstance(msg.content, list) or not any( - isinstance(b, dict) and b.get("type") == "tool_result" for b in msg.content - ): - break - recent_n += 1 - split_index = max(0, len(messages) - max(recent_n, self.recent_n)) - - md_file_tool_ids = set() - try: - for msg in messages: - if not isinstance(msg.content, list): - continue - - for block in msg.content: - if isinstance(block, dict) and block.get("type") == "tool_use": - tool_id = block.get("id", "") - if not tool_id: - continue - - if ( - block.get("name", "").lower() == "read_file" - and ".md" in (block.get("raw_input") or "").lower() - ): - md_file_tool_ids.add(tool_id) - except Exception as e: - logger.warning("Failed to detect md file tool ids: %s", e) - logger.info(f"md_file_tool_ids: {md_file_tool_ids}") - - for idx, msg in enumerate(messages): - if not isinstance(msg.content, list): - continue - is_recent = idx >= split_index - max_bytes = self.recent_max_bytes if is_recent else self.old_max_bytes - for block in msg.content: - if isinstance(block, dict) and block.get("type") == "tool_result" and block.get("output"): - tool_use_id = block.get("id", "") - if tool_use_id in md_file_tool_ids: - effective_max_bytes = self.recent_max_bytes - else: - effective_max_bytes = max_bytes - block["output"] = self._compact(block["output"], effective_max_bytes) - - return messages - - def cleanup_expired_files(self) -> int: - """Clean up files older than retention_days. - - Returns: - Number of files successfully deleted. - """ - if not self.tool_result_dir.exists(): - return 0 - - cutoff = datetime.now() - timedelta(days=self.retention_days) - deleted = failed = 0 - - for fp in self.tool_result_dir.glob("*.txt"): - try: - stat = os.stat(fp) - if sys.platform == "win32": - ts = stat.st_ctime # creation time on Windows - else: - ts = getattr(stat, "st_birthtime", stat.st_mtime) # macOS/BSD; Linux fallback to mtime - if datetime.fromtimestamp(ts) < cutoff: - fp.unlink() - deleted += 1 - except FileNotFoundError: - pass # deleted by another process between glob and stat/unlink - except Exception as e: - failed += 1 - logger.warning("Failed to delete %s: %s", fp, e) - - if deleted or failed: - logger.info("Cleaned up %d expired files (%d failed)", deleted, failed) - return deleted diff --git a/reme/memory/file_based/reme_in_memory_memory.py b/reme/memory/file_based/reme_in_memory_memory.py deleted file mode 100644 index d8293e0e..00000000 --- a/reme/memory/file_based/reme_in_memory_memory.py +++ /dev/null @@ -1,299 +0,0 @@ -"""Custom memory implementation with bugfixes and extensions.""" - -import json -from datetime import datetime -from pathlib import Path - -from agentscope.agent._react_agent import _MemoryMark # noqa -from agentscope.memory import InMemoryMemory -from agentscope.message import Msg -from agentscope.token import HuggingFaceTokenCounter - -from .utils import AsMsgHandler -from ...core.utils import get_logger - -logger = get_logger() - - -class ReMeInMemoryMemory(InMemoryMemory): - """Extended InMemoryMemory with bugfixes and summary support.""" - - def __init__( - self, - token_counter: HuggingFaceTokenCounter, - dialog_path: str | Path | None = None, - ): - """Initialize the ReMeInMemoryMemory. - - Args: - token_counter: Token counter for measuring content length. - dialog_path: Path to the dialog storage directory. If provided, - messages will be persisted to jsonl files when cleared or compressed. - """ - super().__init__() - self._token_counter: HuggingFaceTokenCounter = token_counter - self._msg_handler: AsMsgHandler = AsMsgHandler(token_counter) - self._dialog_path: Path | None = Path(dialog_path) if dialog_path else None - self._long_term_memory: str = "" - - def _append_messages_to_dialog(self, messages: list[Msg]) -> int: - """Append messages to dialog storage file. - - Saves messages to jsonl files named by message date (YYYY-mm-dd.jsonl). - Each line is a JSON representation of a message. - Messages are grouped by their timestamp date. - - Args: - messages: List of messages to append to the dialog file. - - Returns: - Number of messages successfully appended. - """ - if not messages: - return 0 - - if self._dialog_path is None: - logger.warning("dialog_path is not set, skipping dialog persistence") - return 0 - - # Ensure dialog directory exists - try: - self._dialog_path.mkdir(parents=True, exist_ok=True) - except Exception as e: - logger.exception(f"Failed to create dialog directory {self._dialog_path}: {e}") - return 0 - - # Group messages by date (extracted from timestamp) - # timestamp format: "YYYY-mm-dd HH:MM:SS.fff" - messages_by_date: dict[str, list[Msg]] = {} - for msg in messages: - try: - if msg.timestamp: - # Extract date part from timestamp - date_str = msg.timestamp.split()[0] # "YYYY-mm-dd" - else: - date_str = datetime.now().strftime("%Y-%m-%d") - - if date_str not in messages_by_date: - messages_by_date[date_str] = [] - messages_by_date[date_str].append(msg) - except Exception as e: - logger.warning(f"Failed to process message timestamp: {e}, using today's date") - date_str = datetime.now().strftime("%Y-%m-%d") - if date_str not in messages_by_date: - messages_by_date[date_str] = [] - messages_by_date[date_str].append(msg) - - # Append messages to corresponding date files (sorted by timestamp within each date) - total_count = 0 - for date_str, msgs in messages_by_date.items(): - # Sort messages by timestamp within the same date - try: - msgs_sorted = sorted(msgs, key=lambda m: m.timestamp or "") - except Exception as e: - logger.warning(f"Failed to sort messages by timestamp: {e}") - msgs_sorted = msgs - - filename = f"{date_str}.jsonl" - filepath = self._dialog_path / filename - - try: - with open(filepath, "a", encoding="utf-8") as f: - for msg in msgs_sorted: - msg_dict = msg.to_dict() - f.write(json.dumps(msg_dict, ensure_ascii=False) + "\n") - total_count += 1 - logger.info(f"Appended {len(msgs_sorted)} messages to {filepath}") - except Exception as e: - logger.exception(f"Failed to append messages to dialog file {filepath}: {e}") - - return total_count - - async def get_memory( - self, - prepend_summary: bool = True, - **_kwargs, - ) -> list[Msg]: - """Get the messages from the memory by mark (if provided). - - Args: - prepend_summary: Whether to prepend compressed summary - **_kwargs: Additional keyword arguments (ignored) - - Returns: - List of filtered messages - """ - filtered_content = [(msg, marks) for msg, marks in self.content if _MemoryMark.COMPRESSED not in marks] - - parts = [] - if self._long_term_memory: - parts.append(f"# Memories\n\n{self._long_term_memory}") - if prepend_summary and self._compressed_summary: - parts.append( - f"# Summary of previous conversation\n\n" - f"Previous conversation logs are offloaded to dialog/YYYY-MM-DD.jsonl (or nearby date files). " - "Here is the summary:\n\n" - f"{self._compressed_summary}\n" - f"The above is a summary of previous conversation, use it as context to maintain continuity.", - ) - - if parts: - return [Msg("user", "\n\n".join(parts), "user"), *[msg for msg, _ in filtered_content]] - - return [msg for msg, _ in filtered_content] - - def get_compressed_summary(self) -> str: - """Get the compressed summary of the memory.""" - return self._compressed_summary - - def state_dict(self) -> dict: - """Get the state dictionary for serialization.""" - return { - "content": [[msg.to_dict(), marks] for msg, marks in self.content], - "_compressed_summary": self._compressed_summary, - } - - # pylint: disable=attribute-defined-outside-init - def load_state_dict(self, state_dict: dict, strict: bool = True) -> None: - """Load the state dictionary for deserialization.""" - if strict and "content" not in state_dict: - raise KeyError("The state_dict does not contain 'content' key required for InMemoryMemory.") - - self.content = [] # pylint: disable=attribute-defined-outside-init - for item in state_dict.get("content", []): - if isinstance(item, (tuple, list)) and len(item) == 2: - msg_dict, marks = item - msg = Msg.from_dict(msg_dict) - self.content.append((msg, marks)) - - elif isinstance(item, dict): - # For compatibility with older versions - msg = Msg.from_dict(item) - self.content.append((msg, [])) - - else: - raise ValueError("Invalid item format in state_dict for InMemoryMemory.") - - self._compressed_summary = state_dict.get("_compressed_summary", "") - - async def mark_messages_compressed(self, messages: list[Msg]) -> int: - """Mark messages as compressed, persist them to dialog, and remove from memory. - - This method: - 1. Persists the given messages to the dialog storage - 2. Removes them from memory - - Args: - messages: List of messages to mark as compressed. - - Returns: - Number of messages marked as compressed. - """ - if not messages: - return 0 - - # Persist messages to dialog storage instead of compressed - self._append_messages_to_dialog(messages) - - # Remove messages from memory - msg_ids = {msg.id for msg in messages} - initial_size = len(self.content) - self.content = [(msg, marks) for msg, marks in self.content if msg.id not in msg_ids] - removed_count = initial_size - len(self.content) - - logger.info(f"Marked {removed_count} messages as compressed and removed from memory") - return removed_count - - def clear_compressed_summary(self): - """Clear the compressed summary.""" - self._compressed_summary = "" # pylint: disable=attribute-defined-outside-init - - def clear_content(self): - """Persist all messages to dialog storage and clear the content. - - This method: - 1. Persists all messages in memory to the dialog storage - 2. Clears the in-memory content - """ - # Persist all messages to dialog storage - if self.content: - messages = [msg for msg, _ in self.content] - self._append_messages_to_dialog(messages) - - # Clear in-memory content - self.content.clear() - logger.info("Cleared all messages from memory") - - async def estimate_tokens(self, max_input_length: int) -> dict: - """Estimate token usage for current memory. - - Args: - max_input_length: Max input length for context usage calculation. - - Returns: - Dict containing detailed token statistics: - - total_messages: Number of messages - - compressed_summary_tokens: Tokens in compressed summary - - messages_tokens: Tokens in messages - - estimated_tokens: Total estimated tokens - - max_input_length: Max input length from config - - context_usage_ratio: Usage percentage - - messages_detail: List of per-message AsMsgStat objects - """ - messages = await self.get_memory(prepend_summary=False) - - compressed_summary = self.get_compressed_summary() - compressed_summary_tokens = await self._msg_handler.count_str_token(compressed_summary) - - # Build per-message token details using AsMsgHandler - messages_detail = [await self._msg_handler.stat_message(msg) for msg in messages] - - # Calculate total message tokens from stats - messages_tokens = sum(stat.total_tokens for stat in messages_detail) - estimated_tokens = messages_tokens + compressed_summary_tokens - - # Calculate context usage ratio - context_usage_ratio = (estimated_tokens / max_input_length * 100) if max_input_length > 0 else 0 - - return { - "total_messages": len(messages), - "compressed_summary_tokens": compressed_summary_tokens, - "messages_tokens": messages_tokens, - "estimated_tokens": estimated_tokens, - "max_input_length": max_input_length, - "context_usage_ratio": context_usage_ratio, - "messages_detail": messages_detail, - } - - async def get_history_str(self, max_input_length: int) -> str: - """Get formatted history string similar to /history command output. - - Args: - max_input_length: Max input length for context usage calculation. - - Returns: - Formatted string containing conversation history details - """ - stats = await self.estimate_tokens(max_input_length) - - lines = [] - for i, msg_stat in enumerate(stats["messages_detail"], 1): - blocks_info = "" - if msg_stat.content: - block_strs = [f"{b.block_type}(tokens={b.token_count})" for b in msg_stat.content] - blocks_info = f"\n content: [{', '.join(block_strs)}]" - - lines.append( - f"[{i}] **{msg_stat.role}** " - f"(total_tokens={msg_stat.total_tokens})" - f"{blocks_info}\n preview: {msg_stat.preview}", - ) - - return ( - f"**Conversation History**\n\n" - f"- Total messages: {stats['total_messages']}\n" - f"- Estimated tokens: {stats['estimated_tokens']}\n" - f"- Max input length: {stats['max_input_length']}\n" - f"- Context usage: {stats['context_usage_ratio']:.1f}%\n" - f"- Compressed summary tokens: {stats['compressed_summary_tokens']}\n\n" + "\n\n".join(lines) - ) diff --git a/reme/memory/file_based/tools/__init__.py b/reme/memory/file_based/tools/__init__.py deleted file mode 100644 index 0fb0d814..00000000 --- a/reme/memory/file_based/tools/__init__.py +++ /dev/null @@ -1,13 +0,0 @@ -"""File-based memory tool implementations.""" - -from .file_io import FileIO -from .memory_get import MemoryGet -from .memory_search import MemorySearch -from .shell import Shell - -__all__ = [ - "FileIO", - "MemoryGet", - "MemorySearch", - "Shell", -] diff --git a/reme/memory/file_based/tools/browser_control.py b/reme/memory/file_based/tools/browser_control.py deleted file mode 100644 index 8baf9f50..00000000 --- a/reme/memory/file_based/tools/browser_control.py +++ /dev/null @@ -1,2624 +0,0 @@ -# -*- coding: utf-8 -*- -# flake8: noqa: E501 -# pylint: disable=too-many-lines -"""Browser automation tool using Playwright. - -Single tool with action-based API matching browser MCP: start, stop, open, -navigate, navigate_back, screenshot, snapshot, click, type, eval, evaluate, -resize, console_messages, handle_dialog, file_upload, fill_form, install, -press_key, network_requests, run_code, drag, hover, select_option, tabs, -wait_for, pdf, close. Uses refs from snapshot for ref-based actions. -""" - -import asyncio -import atexit -import json -import logging -import os -import subprocess -import sys -import time -from concurrent.futures import ThreadPoolExecutor -from typing import Any, Optional - -from agentscope.message import TextBlock -from agentscope.tool import ToolResponse - -from ...config import ( - get_playwright_chromium_executable_path, - get_system_default_browser, - is_running_in_container, -) - -from .browser_snapshot import build_role_snapshot_from_aria - -logger = logging.getLogger(__name__) - -# Hybrid mode detection: Windows + Uvicorn reload mode requires sync Playwright -# to avoid NotImplementedError with asyncio.create_subprocess_exec. -# On other platforms or without reload, use async Playwright for better performance. -_USE_SYNC_PLAYWRIGHT = sys.platform == "win32" and os.environ.get("COPAW_RELOAD_MODE") == "1" - -if _USE_SYNC_PLAYWRIGHT: - _executor: Optional[ThreadPoolExecutor] = None - - def _get_executor() -> ThreadPoolExecutor: - global _executor - if _executor is None: - _executor = ThreadPoolExecutor( - max_workers=1, - thread_name_prefix="playwright", - ) - return _executor - - async def _run_sync(func, *args, **kwargs): - """Run a sync function in the thread pool and await the result.""" - loop = asyncio.get_event_loop() - return await loop.run_in_executor( - _get_executor(), - lambda: func(*args, **kwargs), - ) - -else: - - async def _run_sync(func, *args, **kwargs): - """Fallback: directly call async function (should not be used in async mode).""" - return await func(*args, **kwargs) - - -# Process-global browser state (one browser, multiple pages by page_id) -_state: dict[str, Any] = { - "playwright": None, - "browser": None, - "context": None, - "pages": {}, - "refs": {}, # page_id -> ref -> {role, name?, nth?} - "refs_frame": {}, # page_id -> frame for last snapshot - "console_logs": {}, # page_id -> list of {level, text} - "network_requests": {}, # page_id -> list of request dicts - "pending_dialogs": {}, # page_id -> dialog handlers - "pending_file_choosers": {}, # page_id -> FileChooser list - "headless": True, - "current_page_id": None, - "page_counter": 0, # monotonic counter for page_N ids, avoids reuse after close - "last_activity_time": 0.0, # monotonic timestamp of last browser activity - "_idle_task": None, # background asyncio.Task for idle watchdog - "_last_browser_error": None, # message when launch failed (for user-facing error) - "_sync_browser": None, # sync browser handle for hybrid mode - "_sync_context": None, # sync context handle for hybrid mode - "_sync_playwright": None, # sync playwright handle for hybrid mode -} - -# Stop the browser after this many seconds of inactivity (default 30 minutes). -_BROWSER_IDLE_TIMEOUT = 1800.0 - - -def _touch_activity() -> None: - """Record the current time as the last browser activity timestamp.""" - _state["last_activity_time"] = time.monotonic() - - -def _is_browser_running() -> bool: - """Check if browser is currently running (sync or async mode).""" - if _USE_SYNC_PLAYWRIGHT: - return _state.get("_sync_browser") is not None - return _state.get("browser") is not None - - -def _reset_browser_state() -> None: - """Reset all browser-related state variables.""" - # Clear sync/async specific state - _state["playwright"] = None - _state["browser"] = None - _state["context"] = None - _state["_sync_playwright"] = None - _state["_sync_browser"] = None - _state["_sync_context"] = None - # Clear shared state - _state["pages"].clear() - _state["refs"].clear() - _state["refs_frame"].clear() - _state["console_logs"].clear() - _state["network_requests"].clear() - _state["pending_dialogs"].clear() - _state["pending_file_choosers"].clear() - _state["current_page_id"] = None - _state["page_counter"] = 0 - _state["last_activity_time"] = 0.0 - _state["headless"] = True - - -async def _idle_watchdog(idle_seconds: float = _BROWSER_IDLE_TIMEOUT) -> None: - """Background task: stop the browser after it has been idle for *idle_seconds*. - - This reclaims Chrome renderer processes that accumulate when pages are - opened during agent tasks but never explicitly closed. - """ - try: - while True: - await asyncio.sleep(60) # check every minute - if not _is_browser_running(): - return - idle = time.monotonic() - _state.get("last_activity_time", 0.0) - if idle >= idle_seconds: - logger.info( - "Browser idle for %.0fs (limit %.0fs), stopping to release resources", - idle, - idle_seconds, - ) - await _action_stop() - return - except asyncio.CancelledError: - pass - - -def _atexit_cleanup() -> None: - """Best-effort browser cleanup registered with :func:`atexit`. - - Playwright child processes are cleaned up by the OS when the parent - exits, but this gives Playwright a chance to flush any pending I/O and - close Chrome gracefully before the process disappears. - """ - if not _is_browser_running(): - return - - try: - loop = asyncio.get_event_loop() - if not loop.is_running() and not loop.is_closed(): - loop.run_until_complete(_action_stop()) - except Exception: - pass - - -atexit.register(_atexit_cleanup) - - -def _tool_response(text: str) -> ToolResponse: - """Wrap text for agentscope Toolkit (return ToolResponse).""" - return ToolResponse( - content=[TextBlock(type="text", text=text)], - ) - - -def _chromium_launch_args() -> list[str]: - """Extra args for Chromium when running in container.""" - if is_running_in_container(): - return ["--no-sandbox", "--disable-dev-shm-usage"] - return [] - - -def _chromium_executable_path() -> str | None: - """Chromium executable path when set (e.g. container); else None.""" - return get_playwright_chromium_executable_path() - - -def _use_webkit_fallback() -> bool: - """True only on macOS when no system Chrome/Edge/Chromium found. - Use WebKit (Safari) to avoid downloading Chromium. Windows has no system - WebKit, so we never use webkit there. - """ - return sys.platform == "darwin" and _chromium_executable_path() is None - - -def _ensure_playwright_async(): - """Import async_playwright; raise ImportError with hint if missing.""" - try: - from playwright.async_api import async_playwright - - return async_playwright - except ImportError as exc: - raise ImportError( - "Playwright not installed. Use the same Python that runs CoPaw (e.g. " - "activate your venv or use 'uv run'): " - f"'{sys.executable}' -m pip install playwright && " - f"'{sys.executable}' -m playwright install", - ) from exc - - -def _ensure_playwright_sync(): - """Import sync_playwright; raise ImportError with hint if missing.""" - try: - from playwright.sync_api import sync_playwright - - return sync_playwright - except ImportError as exc: - raise ImportError( - "Playwright not installed. Use the same Python that runs CoPaw (e.g. " - "activate your venv or use 'uv run'): " - f"'{sys.executable}' -m pip install playwright && " - f"'{sys.executable}' -m playwright install", - ) from exc - - -def _sync_browser_launch(headless: bool): - """Launch browser using sync Playwright (for hybrid mode).""" - sync_playwright = _ensure_playwright_sync() - pw = sync_playwright().start() # Start without context manager - use_default = not is_running_in_container() and os.environ.get( - "COPAW_BROWSER_USE_DEFAULT", - "1", - ).strip().lower() in ("1", "true", "yes") - default_kind, default_path = get_system_default_browser() if use_default else (None, None) - exe: Optional[str] = None - if default_kind == "chromium" and default_path: - exe = default_path - elif default_kind != "webkit": - exe = _chromium_executable_path() - - if exe: - launch_kwargs = {"headless": headless} - extra_args = _chromium_launch_args() - if extra_args: - launch_kwargs["args"] = extra_args - launch_kwargs["executable_path"] = exe - browser = pw.chromium.launch(**launch_kwargs) - elif default_kind == "webkit" or sys.platform == "darwin": - browser = pw.webkit.launch(headless=headless) - else: - launch_kwargs = {"headless": headless} - extra_args = _chromium_launch_args() - if extra_args: - launch_kwargs["args"] = extra_args - browser = pw.chromium.launch(**launch_kwargs) - - context = browser.new_context() - _attach_context_listeners(context) - return pw, browser, context - - -def _sync_browser_close(): - """Close browser using sync Playwright (for hybrid mode).""" - if _state["_sync_browser"] is not None: - try: - _state["_sync_browser"].close() - except Exception: - pass - if _state["_sync_playwright"] is not None: - try: - _state["_sync_playwright"].stop() - except Exception: - pass - - -def _parse_json_param(value: str, default: Any = None): - """Parse optional JSON string param (e.g. fields, paths, values).""" - if not value or not isinstance(value, str): - return default - value = value.strip() - if not value: - return default - try: - return json.loads(value) - except json.JSONDecodeError: - if "," in value: - return [x.strip() for x in value.split(",")] - return default - - -async def browser_use( # pylint: disable=R0911,R0912 - action: str, - url: str = "", - page_id: str = "default", - selector: str = "", - text: str = "", - code: str = "", - path: str = "", - wait: int = 0, - full_page: bool = False, - width: int = 0, - height: int = 0, - level: str = "info", - filename: str = "", - accept: bool = True, - prompt_text: str = "", - ref: str = "", - element: str = "", - paths_json: str = "", - fields_json: str = "", - key: str = "", - submit: bool = False, - slowly: bool = False, - include_static: bool = False, - screenshot_type: str = "png", - snapshot_filename: str = "", - double_click: bool = False, - button: str = "left", - modifiers_json: str = "", - start_ref: str = "", - end_ref: str = "", - start_selector: str = "", - end_selector: str = "", - start_element: str = "", - end_element: str = "", - values_json: str = "", - tab_action: str = "", - index: int = -1, - wait_time: float = 0, - text_gone: str = "", - frame_selector: str = "", - headed: bool = False, -) -> ToolResponse: - """Control browser (Playwright). Default is headless. Use headed=True with - action=start to open a visible browser window. Flow: start, open(url), - snapshot to get refs, then click/type etc. with ref or selector. Use - page_id for multiple tabs. - - Args: - action (str): - Required. Action type. Values: start, stop, open, navigate, - navigate_back, snapshot, screenshot, click, type, eval, evaluate, - resize, console_messages, network_requests, handle_dialog, - file_upload, fill_form, install, press_key, run_code, drag, hover, - select_option, tabs, wait_for, pdf, close. - url (str): - URL to open. Required for action=open or navigate. - page_id (str): - Page/tab identifier, default "default". Use different page_id for - multiple tabs. - selector (str): - CSS selector to locate element for click/type/hover etc. Prefer - ref when available. - text (str): - Text to type. Required for action=type. - code (str): - JavaScript code. Required for action=eval, evaluate, or run_code. - path (str): - File path for screenshot save or PDF export. - wait (int): - Milliseconds to wait after click. Used with action=click. - full_page (bool): - Whether to capture full page. Used with action=screenshot. - width (int): - Viewport width in pixels. Used with action=resize. - height (int): - Viewport height in pixels. Used with action=resize. - level (str): - Console log level filter, e.g. "info" or "error". Used with - action=console_messages. - filename (str): - Filename for saving logs or screenshot. Used with - console_messages, network_requests, screenshot. - accept (bool): - Whether to accept dialog (true) or dismiss (false). Used with - action=handle_dialog. - prompt_text (str): - Input for prompt dialog. Used with action=handle_dialog when - dialog is prompt. - ref (str): - Element ref from snapshot output; use for stable targeting. Prefer - ref for click/type/hover/screenshot/evaluate/select_option. - element (str): - Element description for evaluate etc. Prefer ref when available. - paths_json (str): - JSON array string of file paths. Used with action=file_upload. - fields_json (str): - JSON object string of form field name to value. Used with - action=fill_form. - key (str): - Key name, e.g. "Enter", "Control+a". Required for - action=press_key. - submit (bool): - Whether to submit (press Enter) after typing. Used with - action=type. - slowly (bool): - Whether to type character by character. Used with action=type. - include_static (bool): - Whether to include static resource requests. Used with - action=network_requests. - screenshot_type (str): - Screenshot format, "png" or "jpeg". Used with action=screenshot. - snapshot_filename (str): - File path to save snapshot output. Used with action=snapshot. - double_click (bool): - Whether to double-click. Used with action=click. - button (str): - Mouse button: "left", "right", or "middle". Used with - action=click. - modifiers_json (str): - JSON array of modifier keys, e.g. ["Shift","Control"]. Used with - action=click. - start_ref (str): - Drag start element ref. Used with action=drag. - end_ref (str): - Drag end element ref. Used with action=drag. - start_selector (str): - Drag start CSS selector. Used with action=drag. - end_selector (str): - Drag end CSS selector. Used with action=drag. - start_element (str): - Drag start element description. Used with action=drag. - end_element (str): - Drag end element description. Used with action=drag. - values_json (str): - JSON of option value(s) for select. Used with - action=select_option. - tab_action (str): - Tab action: list, new, close, or select. Required for - action=tabs. - index (int): - Tab index for tabs select, zero-based. Used with action=tabs. - wait_time (float): - Seconds to wait. Used with action=wait_for. - text_gone (str): - Wait until this text disappears from page. Used with - action=wait_for. - frame_selector (str): - iframe selector, e.g. "iframe#main". Set when operating inside - that iframe in snapshot/click/type etc. - headed (bool): - When True with action=start, launch a visible browser window - (non-headless). User can see the real browser. Default False. - """ - action = (action or "").strip().lower() - if not action: - return _tool_response( - json.dumps( - {"ok": False, "error": "action required"}, - ensure_ascii=False, - indent=2, - ), - ) - - page_id = (page_id or "default").strip() or "default" - current = _state.get("current_page_id") - pages = _state.get("pages") or {} - if page_id == "default" and current and current in pages: - page_id = current - - try: - if action == "start": - return await _action_start(headed=headed) - if action == "stop": - return await _action_stop() - if action == "open": - return await _action_open(url, page_id) - if action == "navigate": - return await _action_navigate(url, page_id) - if action == "navigate_back": - return await _action_navigate_back(page_id) - if action in ("screenshot", "take_screenshot"): - return await _action_screenshot( - page_id, - path or filename, - full_page, - screenshot_type, - ref, - element, - frame_selector, - ) - if action == "snapshot": - return await _action_snapshot( - page_id, - snapshot_filename or filename, - frame_selector, - ) - if action == "click": - return await _action_click( - page_id, - selector, - ref, - element, - wait, - double_click, - button, - modifiers_json, - frame_selector, - ) - if action == "type": - return await _action_type( - page_id, - selector, - ref, - element, - text, - submit, - slowly, - frame_selector, - ) - if action == "eval": - return await _action_eval(page_id, code) - if action == "evaluate": - return await _action_evaluate( - page_id, - code, - ref, - element, - frame_selector, - ) - if action == "resize": - return await _action_resize(page_id, width, height) - if action == "console_messages": - return await _action_console_messages( - page_id, - level, - filename or path, - ) - if action == "handle_dialog": - return await _action_handle_dialog(page_id, accept, prompt_text) - if action == "file_upload": - return await _action_file_upload(page_id, paths_json) - if action == "fill_form": - return await _action_fill_form(page_id, fields_json) - if action == "install": - return await _action_install() - if action == "press_key": - return await _action_press_key(page_id, key) - if action == "network_requests": - return await _action_network_requests( - page_id, - include_static, - filename or path, - ) - if action == "run_code": - return await _action_run_code(page_id, code) - if action == "drag": - return await _action_drag( - page_id, - start_ref, - end_ref, - start_selector, - end_selector, - start_element, - end_element, - frame_selector, - ) - if action == "hover": - return await _action_hover( - page_id, - ref, - element, - selector, - frame_selector, - ) - if action == "select_option": - return await _action_select_option( - page_id, - ref, - element, - values_json, - frame_selector, - ) - if action == "tabs": - return await _action_tabs(page_id, tab_action, index) - if action == "wait_for": - return await _action_wait_for(page_id, wait_time, text, text_gone) - if action == "pdf": - return await _action_pdf(page_id, path) - if action == "close": - return await _action_close(page_id) - return _tool_response( - json.dumps( - {"ok": False, "error": f"Unknown action: {action}"}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - logger.exception("Browser tool error: %s", e, exc_info=True) - return _tool_response( - json.dumps( - {"ok": False, "error": str(e)}, - ensure_ascii=False, - indent=2, - ), - ) - - -def _get_page(page_id: str): - """Return page for page_id or None if not found.""" - return _state["pages"].get(page_id) - - -def _get_refs(page_id: str) -> dict[str, dict]: - """Return refs map for page_id (ref -> {role, name?, nth?}).""" - return _state["refs"].setdefault(page_id, {}) - - -def _get_root(page, _page_id: str, frame_selector: str = ""): - """Return page or frame for frame_selector (ref/selector).""" - if not (frame_selector and frame_selector.strip()): - return page - return page.frame_locator(frame_selector.strip()) - - -def _get_locator_by_ref( - page, - page_id: str, - ref: str, - frame_selector: str = "", -): - """Resolve snapshot ref to locator; frame_selector for iframe.""" - refs = _get_refs(page_id) - info = refs.get(ref) - if not info: - return None - role = info.get("role", "generic") - name = info.get("name") - nth = info.get("nth", 0) - root = _get_root(page, page_id, frame_selector) - locator = root.get_by_role(role, name=name or None) - if nth is not None and nth > 0: - locator = locator.nth(nth) - return locator - - -def _attach_page_listeners(page, page_id: str) -> None: - """Attach console and request listeners for a page.""" - logs = _state["console_logs"].setdefault(page_id, []) - - def on_console(msg): - logs.append({"level": msg.type, "text": msg.text}) - - page.on("console", on_console) - requests_list = _state["network_requests"].setdefault(page_id, []) - - def on_request(req): - requests_list.append( - { - "url": req.url, - "method": req.method, - "resourceType": getattr(req, "resource_type", None), - }, - ) - - def on_response(res): - for r in requests_list: - if r.get("url") == res.url and "status" not in r: - r["status"] = res.status - break - - page.on("request", on_request) - page.on("response", on_response) - dialogs = _state["pending_dialogs"].setdefault(page_id, []) - - def on_dialog(dialog): - dialogs.append(dialog) - - page.on("dialog", on_dialog) - choosers = _state["pending_file_choosers"].setdefault(page_id, []) - - def on_filechooser(chooser): - choosers.append(chooser) - - page.on("filechooser", on_filechooser) - - -def _next_page_id() -> str: - """Return a unique page_id (page_N). - Uses monotonic counter so IDs are not reused after close.""" - _state["page_counter"] = _state.get("page_counter", 0) + 1 - return f"page_{_state['page_counter']}" - - -def _attach_context_listeners(context) -> None: - """When the page opens a new tab (e.g. target=_blank, window.open), - register it and set as current.""" - - def on_page(page): - new_id = _next_page_id() - _state["refs"][new_id] = {} - _state["console_logs"][new_id] = [] - _state["network_requests"][new_id] = [] - _state["pending_dialogs"][new_id] = [] - _state["pending_file_choosers"][new_id] = [] - _attach_page_listeners(page, new_id) - _state["pages"][new_id] = page - _state["current_page_id"] = new_id - logger.debug( - "New tab opened by page, registered as page_id=%s", - new_id, - ) - - context.on("page", on_page) - - -async def _ensure_browser() -> bool: # pylint: disable=too-many-branches - """Start browser if not running. Return True if ready, False on failure.""" - # Check browser state based on mode - if _USE_SYNC_PLAYWRIGHT: - if _state["_sync_browser"] is not None and _state["_sync_context"] is not None: - _touch_activity() - return True - else: - if _state["browser"] is not None and _state["context"] is not None: - _touch_activity() - return True - - try: - if _USE_SYNC_PLAYWRIGHT: - # Hybrid mode: use sync Playwright in thread pool - loop = asyncio.get_event_loop() - pw, browser, context = await loop.run_in_executor( - _get_executor(), - lambda: _sync_browser_launch(_state["headless"]), - ) - _state["_sync_playwright"] = pw - _state["_sync_browser"] = browser - _state["_sync_context"] = context - else: - # Standard mode: use async Playwright - async_playwright = _ensure_playwright_async() - pw = await async_playwright().start() - # Prefer OS default browser when available (e.g. user's default Chrome/Safari). - use_default = not is_running_in_container() and os.environ.get( - "COPAW_BROWSER_USE_DEFAULT", - "1", - ).strip().lower() in ("1", "true", "yes") - default_kind, default_path = get_system_default_browser() if use_default else (None, None) - exe: Optional[str] = None - if default_kind == "chromium" and default_path: - exe = default_path - elif default_kind != "webkit": - exe = _chromium_executable_path() - if exe: - # System Chrome/Edge/Chromium (default or discovered) - launch_kwargs: dict[str, Any] = { - "headless": _state["headless"], - } - extra_args = _chromium_launch_args() - if extra_args: - launch_kwargs["args"] = extra_args - launch_kwargs["executable_path"] = exe - pw_browser = await pw.chromium.launch(**launch_kwargs) - elif default_kind == "webkit" or sys.platform == "darwin": - # macOS: default Safari or no Chromium → use WebKit (Safari) - pw_browser = await pw.webkit.launch( - headless=_state["headless"], - ) - else: - # Windows/Linux without system Chromium → Playwright's Chromium - launch_kwargs = {"headless": _state["headless"]} - extra_args = _chromium_launch_args() - if extra_args: - launch_kwargs["args"] = extra_args - pw_browser = await pw.chromium.launch(**launch_kwargs) - context = await pw_browser.new_context() - _attach_context_listeners(context) - _state["playwright"] = pw - _state["browser"] = pw_browser - _state["context"] = context - _state["_last_browser_error"] = None - _touch_activity() - _start_idle_watchdog() - return True - except Exception as e: - _state["_last_browser_error"] = str(e) - return False - - -def _start_idle_watchdog() -> None: - """Cancel any existing idle watchdog and start a fresh one.""" - old_task = _state.get("_idle_task") - if old_task and not old_task.done(): - old_task.cancel() - _state["_idle_task"] = asyncio.ensure_future(_idle_watchdog()) - - -def _cancel_idle_watchdog() -> None: - """Cancel the idle watchdog, if running.""" - task = _state.get("_idle_task") - if task and not task.done(): - task.cancel() - _state["_idle_task"] = None - - -# pylint: disable=R0912,R0915 -async def _action_start( - headed: bool = False, -) -> ToolResponse: - # Check browser state based on mode - if _USE_SYNC_PLAYWRIGHT: - browser_exists = _state["_sync_browser"] is not None - current_headless = not _state.get("_sync_headless", True) - else: - browser_exists = _state["browser"] is not None - current_headless = _state["headless"] - - # If user asks for visible window (headed=True) - # but browser is already running headless, restart with headed - if browser_exists: - if headed and current_headless: - _cancel_idle_watchdog() - try: - await _action_stop() - except Exception: - pass - else: - return _tool_response( - json.dumps( - {"ok": True, "message": "Browser already running"}, - ensure_ascii=False, - indent=2, - ), - ) - # Default: headless (background). Only headed=True (e.g. browser_visible skill) shows window. - _state["headless"] = not headed - - try: - if _USE_SYNC_PLAYWRIGHT: - loop = asyncio.get_event_loop() - pw, browser, context = await loop.run_in_executor( - _get_executor(), - lambda: _sync_browser_launch(_state["headless"]), - ) - _state["_sync_playwright"] = pw - _state["_sync_browser"] = browser - _state["_sync_context"] = context - _state["_sync_headless"] = not headed - else: - async_playwright = _ensure_playwright_async() - pw = await async_playwright().start() - use_default = not is_running_in_container() and os.environ.get( - "COPAW_BROWSER_USE_DEFAULT", - "1", - ).strip().lower() in ("1", "true", "yes") - default_kind, default_path = get_system_default_browser() if use_default else (None, None) - exe: Optional[str] = None - if default_kind == "chromium" and default_path: - exe = default_path - elif default_kind != "webkit": - exe = _chromium_executable_path() - if exe: - launch_kwargs = {"headless": _state["headless"]} - extra_args = _chromium_launch_args() - if extra_args: - launch_kwargs["args"] = extra_args - launch_kwargs["executable_path"] = exe - pw_browser = await pw.chromium.launch(**launch_kwargs) - elif default_kind == "webkit" or sys.platform == "darwin": - pw_browser = await pw.webkit.launch( - headless=_state["headless"], - ) - else: - launch_kwargs = {"headless": _state["headless"]} - extra_args = _chromium_launch_args() - if extra_args: - launch_kwargs["args"] = extra_args - pw_browser = await pw.chromium.launch(**launch_kwargs) - context = await pw_browser.new_context() - _attach_context_listeners(context) - _state["playwright"] = pw - _state["browser"] = pw_browser - _state["context"] = context - _touch_activity() - _start_idle_watchdog() - msg = "Browser started (visible window)" if not _state["headless"] else "Browser started" - return _tool_response( - json.dumps( - {"ok": True, "message": msg}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Browser start failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_stop() -> ToolResponse: - _cancel_idle_watchdog() - - # Check browser state based on mode - if not _is_browser_running(): - return _tool_response( - json.dumps( - {"ok": True, "message": "Browser not running"}, - ensure_ascii=False, - indent=2, - ), - ) - - if _USE_SYNC_PLAYWRIGHT: - # Close sync browser in thread pool - loop = asyncio.get_event_loop() - try: - await loop.run_in_executor( - _get_executor(), - _sync_browser_close, - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Browser stop failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - finally: - _reset_browser_state() - else: - # Standard async mode - try: - await _state["browser"].close() - if _state["playwright"] is not None: - await _state["playwright"].stop() - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Browser stop failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - finally: - _reset_browser_state() - - return _tool_response( - json.dumps( - {"ok": True, "message": "Browser stopped"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_open(url: str, page_id: str) -> ToolResponse: - url = (url or "").strip() - if not url: - return _tool_response( - json.dumps( - {"ok": False, "error": "url required for open"}, - ensure_ascii=False, - indent=2, - ), - ) - if not await _ensure_browser(): - err = _state.get("_last_browser_error") or "Browser not started" - return _tool_response( - json.dumps( - {"ok": False, "error": err}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if _USE_SYNC_PLAYWRIGHT: - # Hybrid mode: create page in thread pool - loop = asyncio.get_event_loop() - # pylint: disable=unnecessary-lambda - page = await loop.run_in_executor( - _get_executor(), - lambda: _state["_sync_context"].new_page(), - ) - else: - # Standard async mode - page = await _state["context"].new_page() - - _state["refs"][page_id] = {} - _state["console_logs"][page_id] = [] - _state["network_requests"][page_id] = [] - _state["pending_dialogs"][page_id] = [] - _state["pending_file_choosers"][page_id] = [] - _attach_page_listeners(page, page_id) - - if _USE_SYNC_PLAYWRIGHT: - loop = asyncio.get_event_loop() - await loop.run_in_executor( - _get_executor(), - lambda: page.goto(url), - ) - else: - await page.goto(url) - - _state["pages"][page_id] = page - _state["current_page_id"] = page_id - return _tool_response( - json.dumps( - { - "ok": True, - "message": f"Opened {url}", - "page_id": page_id, - "url": url, - }, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Open failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_navigate(url: str, page_id: str) -> ToolResponse: - url = (url or "").strip() - if not url: - return _tool_response( - json.dumps( - {"ok": False, "error": "url required for navigate"}, - ensure_ascii=False, - indent=2, - ), - ) - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if _USE_SYNC_PLAYWRIGHT: - loop = asyncio.get_event_loop() - await loop.run_in_executor( - _get_executor(), - lambda: page.goto(url), - ) - else: - await page.goto(url) - _state["current_page_id"] = page_id - return _tool_response( - json.dumps( - { - "ok": True, - "message": f"Navigated to {url}", - "url": page.url, - }, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Navigate failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_screenshot( - page_id: str, - path: str, - full_page: bool, - screenshot_type: str = "png", - ref: str = "", - element: str = "", # pylint: disable=unused-argument - frame_selector: str = "", -) -> ToolResponse: - path = (path or "").strip() - if not path: - ext = "jpeg" if screenshot_type == "jpeg" else "png" - path = f"page-{int(time.time())}.{ext}" - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if ref and ref.strip(): - locator = _get_locator_by_ref( - page, - page_id, - ref.strip(), - frame_selector, - ) - if locator is None: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Unknown ref: {ref}"}, - ensure_ascii=False, - indent=2, - ), - ) - if _USE_SYNC_PLAYWRIGHT: - await _run_sync( - locator.screenshot, - path=path, - type=screenshot_type if screenshot_type == "jpeg" else "png", - ) - else: - await locator.screenshot( - path=path, - type=screenshot_type if screenshot_type == "jpeg" else "png", - ) - else: - if frame_selector and frame_selector.strip(): - root = _get_root(page, page_id, frame_selector) - locator = root.locator("body").first - if _USE_SYNC_PLAYWRIGHT: - await _run_sync( - locator.screenshot, - path=path, - type=screenshot_type if screenshot_type == "jpeg" else "png", - ) - else: - await locator.screenshot( - path=path, - type=screenshot_type if screenshot_type == "jpeg" else "png", - ) - else: - if _USE_SYNC_PLAYWRIGHT: - await _run_sync( - page.screenshot, - path=path, - full_page=full_page, - type=screenshot_type if screenshot_type == "jpeg" else "png", - ) - else: - await page.screenshot( - path=path, - full_page=full_page, - type=screenshot_type if screenshot_type == "jpeg" else "png", - ) - return _tool_response( - json.dumps( - { - "ok": True, - "message": f"Screenshot saved to {path}", - "path": path, - }, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Screenshot failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_click( # pylint: disable=too-many-branches - page_id: str, - selector: str, - ref: str = "", - element: str = "", # pylint: disable=unused-argument - wait: int = 0, - double_click: bool = False, - button: str = "left", - modifiers_json: str = "", - frame_selector: str = "", -) -> ToolResponse: - ref = (ref or "").strip() - selector = (selector or "").strip() - if not ref and not selector: - return _tool_response( - json.dumps( - {"ok": False, "error": "selector or ref required for click"}, - ensure_ascii=False, - indent=2, - ), - ) - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if wait > 0: - await asyncio.sleep(wait / 1000.0) - mods = _parse_json_param(modifiers_json, []) - if not isinstance(mods, list): - mods = [] - kwargs = { - "button": button if button in ("left", "right", "middle") else "left", - } - if mods: - kwargs["modifiers"] = [m for m in mods if m in ("Alt", "Control", "ControlOrMeta", "Meta", "Shift")] - - if _USE_SYNC_PLAYWRIGHT: - loop = asyncio.get_event_loop() - if ref: - locator = _get_locator_by_ref( - page, - page_id, - ref, - frame_selector, - ) - if locator is None: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Unknown ref: {ref}"}, - ensure_ascii=False, - indent=2, - ), - ) - if double_click: - await loop.run_in_executor( - _get_executor(), - lambda: locator.dblclick(**kwargs), - ) - else: - await loop.run_in_executor( - _get_executor(), - lambda: locator.click(**kwargs), - ) - else: - root = _get_root(page, page_id, frame_selector) - locator = root.locator(selector).first - if double_click: - await loop.run_in_executor( - _get_executor(), - lambda: locator.dblclick(**kwargs), - ) - else: - await loop.run_in_executor( - _get_executor(), - lambda: locator.click(**kwargs), - ) - else: - # Standard async mode - if ref: - locator = _get_locator_by_ref( - page, - page_id, - ref, - frame_selector, - ) - if locator is None: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Unknown ref: {ref}"}, - ensure_ascii=False, - indent=2, - ), - ) - if double_click: - await locator.dblclick(**kwargs) - else: - await locator.click(**kwargs) - else: - root = _get_root(page, page_id, frame_selector) - locator = root.locator(selector).first - if double_click: - await locator.dblclick(**kwargs) - else: - await locator.click(**kwargs) - - return _tool_response( - json.dumps( - {"ok": True, "message": f"Clicked {ref or selector}"}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Click failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_type( - page_id: str, - selector: str, - ref: str = "", - element: str = "", # pylint: disable=unused-argument - text: str = "", - submit: bool = False, - slowly: bool = False, - frame_selector: str = "", -) -> ToolResponse: - ref = (ref or "").strip() - selector = (selector or "").strip() - if not ref and not selector: - return _tool_response( - json.dumps( - {"ok": False, "error": "selector or ref required for type"}, - ensure_ascii=False, - indent=2, - ), - ) - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if ref: - locator = _get_locator_by_ref(page, page_id, ref, frame_selector) - if locator is None: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Unknown ref: {ref}"}, - ensure_ascii=False, - indent=2, - ), - ) - if _USE_SYNC_PLAYWRIGHT: - loop = asyncio.get_event_loop() - if slowly: - await loop.run_in_executor( - _get_executor(), - lambda: locator.press_sequentially(text or ""), - ) - else: - await loop.run_in_executor( - _get_executor(), - lambda: locator.fill(text or ""), - ) - if submit: - await loop.run_in_executor( - _get_executor(), - lambda: locator.press("Enter"), - ) - else: - if slowly: - await locator.press_sequentially(text or "") - else: - await locator.fill(text or "") - if submit: - await locator.press("Enter") - else: - root = _get_root(page, page_id, frame_selector) - loc = root.locator(selector).first - if _USE_SYNC_PLAYWRIGHT: - loop = asyncio.get_event_loop() - if slowly: - await loop.run_in_executor( - _get_executor(), - lambda: loc.press_sequentially(text or ""), - ) - else: - await loop.run_in_executor( - _get_executor(), - lambda: loc.fill(text or ""), - ) - if submit: - await loop.run_in_executor( - _get_executor(), - lambda: loc.press("Enter"), - ) - else: - if slowly: - await loc.press_sequentially(text or "") - else: - await loc.fill(text or "") - if submit: - await loc.press("Enter") - return _tool_response( - json.dumps( - {"ok": True, "message": f"Typed into {ref or selector}"}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Type failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_eval(page_id: str, code: str) -> ToolResponse: - code = (code or "").strip() - if not code: - return _tool_response( - json.dumps( - {"ok": False, "error": "code required for eval"}, - ensure_ascii=False, - indent=2, - ), - ) - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if code.strip().startswith("(") or code.strip().startswith("function"): - if _USE_SYNC_PLAYWRIGHT: - result = await _run_sync(page.evaluate, code) - else: - result = await page.evaluate(code) - else: - if _USE_SYNC_PLAYWRIGHT: - result = await _run_sync( - page.evaluate, - f"() => {{ return ({code}); }}", - ) - else: - result = await page.evaluate(f"() => {{ return ({code}); }}") - try: - out = json.dumps( - {"ok": True, "result": result}, - ensure_ascii=False, - indent=2, - ) - except TypeError: - out = json.dumps( - {"ok": True, "result": str(result)}, - ensure_ascii=False, - indent=2, - ) - return _tool_response(out) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Eval failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_pdf(page_id: str, path: str) -> ToolResponse: - path = (path or "page.pdf").strip() or "page.pdf" - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(page.pdf, path=path) - else: - await page.pdf(path=path) - return _tool_response( - json.dumps( - {"ok": True, "message": f"PDF saved to {path}", "path": path}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"PDF failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_close(page_id: str) -> ToolResponse: - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(page.close) - else: - await page.close() - del _state["pages"][page_id] - for key in ( - "refs", - "refs_frame", - "console_logs", - "network_requests", - "pending_dialogs", - "pending_file_choosers", - ): - _state[key].pop(page_id, None) - if _state.get("current_page_id") == page_id: - remaining = list(_state["pages"].keys()) - _state["current_page_id"] = remaining[0] if remaining else None - return _tool_response( - json.dumps( - {"ok": True, "message": f"Closed page '{page_id}'"}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Close failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_snapshot( - page_id: str, - filename: str, - frame_selector: str = "", -) -> ToolResponse: - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if _USE_SYNC_PLAYWRIGHT: - # Hybrid mode: execute in thread pool - loop = asyncio.get_event_loop() - root = _get_root(page, page_id, frame_selector) - locator = root.locator(":root") - raw = await loop.run_in_executor( - _get_executor(), - lambda: locator.aria_snapshot(), # pylint: disable=unnecessary-lambda - ) - else: - root = _get_root(page, page_id, frame_selector) - locator = root.locator(":root") - raw = await locator.aria_snapshot() - - raw_str = str(raw) if raw is not None else "" - snapshot, refs = build_role_snapshot_from_aria( - raw_str, - interactive=False, - compact=False, - ) - _state["refs"][page_id] = refs - _state["refs_frame"][page_id] = frame_selector.strip() if frame_selector else "" - out = { - "ok": True, - "snapshot": snapshot, - "refs": list(refs.keys()), - "url": page.url, - } - if frame_selector and frame_selector.strip(): - out["frame_selector"] = frame_selector.strip() - if filename and filename.strip(): - with open(filename.strip(), "w", encoding="utf-8") as f: - f.write(snapshot) - out["filename"] = filename.strip() - return _tool_response(json.dumps(out, ensure_ascii=False, indent=2)) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Snapshot failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_navigate_back(page_id: str) -> ToolResponse: - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(page.go_back) - else: - await page.go_back() - return _tool_response( - json.dumps( - {"ok": True, "message": "Navigated back", "url": page.url}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Navigate back failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_evaluate( - page_id: str, - code: str, - ref: str = "", - element: str = "", # pylint: disable=unused-argument - frame_selector: str = "", -) -> ToolResponse: - code = (code or "").strip() - if not code: - return _tool_response( - json.dumps( - {"ok": False, "error": "code required for evaluate"}, - ensure_ascii=False, - indent=2, - ), - ) - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if ref and ref.strip(): - locator = _get_locator_by_ref( - page, - page_id, - ref.strip(), - frame_selector, - ) - if locator is None: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Unknown ref: {ref}"}, - ensure_ascii=False, - indent=2, - ), - ) - if _USE_SYNC_PLAYWRIGHT: - result = await _run_sync(locator.evaluate, code) - else: - result = await locator.evaluate(code) - else: - if code.strip().startswith("(") or code.strip().startswith( - "function", - ): - if _USE_SYNC_PLAYWRIGHT: - result = await _run_sync(page.evaluate, code) - else: - result = await page.evaluate(code) - else: - if _USE_SYNC_PLAYWRIGHT: - result = await _run_sync( - page.evaluate, - f"() => {{ return ({code}); }}", - ) - else: - result = await page.evaluate( - f"() => {{ return ({code}); }}", - ) - try: - out = json.dumps( - {"ok": True, "result": result}, - ensure_ascii=False, - indent=2, - ) - except TypeError: - out = json.dumps( - {"ok": True, "result": str(result)}, - ensure_ascii=False, - indent=2, - ) - return _tool_response(out) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Evaluate failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_resize( - page_id: str, - width: int, - height: int, -) -> ToolResponse: - if width <= 0 or height <= 0: - return _tool_response( - json.dumps( - {"ok": False, "error": "width and height must be positive"}, - ensure_ascii=False, - indent=2, - ), - ) - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if _USE_SYNC_PLAYWRIGHT: - await _run_sync( - page.set_viewport_size, - {"width": width, "height": height}, - ) - else: - await page.set_viewport_size({"width": width, "height": height}) - return _tool_response( - json.dumps( - {"ok": True, "message": f"Resized to {width}x{height}"}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Resize failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_console_messages( - page_id: str, - level: str, - filename: str, -) -> ToolResponse: - level = (level or "info").strip().lower() - order = ("error", "warning", "info", "debug") - idx = order.index(level) if level in order else 2 - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - logs = _state["console_logs"].get(page_id, []) - filtered = [m for m in logs if order.index(m["level"]) <= idx] if level in order else logs - lines = [f"[{m['level']}] {m['text']}" for m in filtered] - text = "\n".join(lines) - if filename and filename.strip(): - with open(filename.strip(), "w", encoding="utf-8") as f: - f.write(text) - return _tool_response( - json.dumps( - { - "ok": True, - "message": f"Console messages saved to {filename}", - "filename": filename.strip(), - }, - ensure_ascii=False, - indent=2, - ), - ) - return _tool_response( - json.dumps( - {"ok": True, "messages": filtered, "text": text}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_handle_dialog( - page_id: str, - accept: bool, - prompt_text: str, -) -> ToolResponse: - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - dialogs = _state["pending_dialogs"].get(page_id, []) - if not dialogs: - return _tool_response( - json.dumps( - {"ok": False, "error": "No pending dialog"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - dialog = dialogs.pop(0) - if accept: - if prompt_text and hasattr(dialog, "accept"): - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(dialog.accept, prompt_text) - else: - await dialog.accept(prompt_text) - else: - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(dialog.accept) - else: - await dialog.accept() - else: - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(dialog.dismiss) - else: - await dialog.dismiss() - return _tool_response( - json.dumps( - {"ok": True, "message": "Dialog handled"}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Handle dialog failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_file_upload(page_id: str, paths_json: str) -> ToolResponse: - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - paths = _parse_json_param(paths_json, []) - if not isinstance(paths, list): - paths = [] - try: - choosers = _state["pending_file_choosers"].get(page_id, []) - if not choosers: - return _tool_response( - json.dumps( - { - "ok": False, - "error": "No chooser. Click upload then file_upload.", - }, - ensure_ascii=False, - indent=2, - ), - ) - chooser = choosers.pop(0) - if paths: - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(chooser.set_files, paths) - else: - await chooser.set_files(paths) - return _tool_response( - json.dumps( - {"ok": True, "message": f"Uploaded {len(paths)} file(s)"}, - ensure_ascii=False, - indent=2, - ), - ) - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(chooser.set_files, []) - else: - await chooser.set_files([]) - return _tool_response( - json.dumps( - {"ok": True, "message": "File chooser cancelled"}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"File upload failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_fill_form(page_id: str, fields_json: str) -> ToolResponse: - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - fields = _parse_json_param(fields_json, []) - if not isinstance(fields, list) or not fields: - return _tool_response( - json.dumps( - {"ok": False, "error": "fields required (JSON array)"}, - ensure_ascii=False, - indent=2, - ), - ) - refs = _get_refs(page_id) - # Use last snapshot's frame so fill_form works after iframe snapshot - frame = _state["refs_frame"].get(page_id, "") - try: - for f in fields: - ref = (f.get("ref") or "").strip() - if not ref or ref not in refs: - continue - locator = _get_locator_by_ref(page, page_id, ref, frame) - if locator is None: - continue - field_type = (f.get("type") or "textbox").lower() - value = f.get("value") - if field_type == "checkbox": - if isinstance(value, str): - value = value.strip().lower() in ("true", "1", "yes") - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(locator.set_checked, bool(value)) - else: - await locator.set_checked(bool(value)) - elif field_type == "radio": - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(locator.set_checked, True) - else: - await locator.set_checked(True) - elif field_type == "combobox": - if _USE_SYNC_PLAYWRIGHT: - await _run_sync( - locator.select_option, - label=value if isinstance(value, str) else None, - value=value, - ) - else: - await locator.select_option( - label=value if isinstance(value, str) else None, - value=value, - ) - elif field_type == "slider": - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(locator.fill, str(value)) - else: - await locator.fill(str(value)) - else: - if _USE_SYNC_PLAYWRIGHT: - await _run_sync( - locator.fill, - str(value) if value is not None else "", - ) - else: - await locator.fill(str(value) if value is not None else "") - return _tool_response( - json.dumps( - {"ok": True, "message": f"Filled {len(fields)} field(s)"}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Fill form failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -def _run_playwright_install() -> None: - """Run playwright install in a blocking way (for use in thread).""" - subprocess.run( - [sys.executable, "-m", "playwright", "install"], - check=True, - capture_output=True, - text=True, - timeout=600, # 10 minutes max - ) - - -async def _action_install() -> ToolResponse: - """Install Playwright browsers. If a system Chrome/Chromium/Edge is found, - use it and skip download. On macOS with no Chromium, use Safari (WebKit) - so no download is needed. Only run playwright install when necessary. - """ - exe = _chromium_executable_path() - if exe: - return _tool_response( - json.dumps( - { - "ok": True, - "message": f"Using system browser (no download): {exe}", - }, - ensure_ascii=False, - indent=2, - ), - ) - if _use_webkit_fallback(): - return _tool_response( - json.dumps( - { - "ok": True, - "message": "On macOS using Safari (WebKit); no browser download needed.", - }, - ensure_ascii=False, - indent=2, - ), - ) - try: - await asyncio.to_thread(_run_playwright_install) - return _tool_response( - json.dumps( - {"ok": True, "message": "Browser installed"}, - ensure_ascii=False, - indent=2, - ), - ) - except subprocess.TimeoutExpired: - return _tool_response( - json.dumps( - { - "ok": False, - "error": "Browser install timed out (10 min). Run manually in terminal: " - f"{sys.executable!s} -m playwright install", - }, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - { - "ok": False, - "error": f"Install failed: {e!s}. Install manually: " - f"{sys.executable!s} -m pip install playwright && " - f"{sys.executable!s} -m playwright install", - }, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_press_key(page_id: str, key: str) -> ToolResponse: - key = (key or "").strip() - if not key: - return _tool_response( - json.dumps( - {"ok": False, "error": "key required for press_key"}, - ensure_ascii=False, - indent=2, - ), - ) - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(page.keyboard.press, key) - else: - await page.keyboard.press(key) - return _tool_response( - json.dumps( - {"ok": True, "message": f"Pressed key {key}"}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Press key failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_network_requests( - page_id: str, - include_static: bool, - filename: str, -) -> ToolResponse: - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - requests = _state["network_requests"].get(page_id, []) - if not include_static: - static = ("image", "stylesheet", "font", "media") - requests = [r for r in requests if r.get("resourceType") not in static] - lines = [f"{r.get('method', '')} {r.get('url', '')} {r.get('status', '')}" for r in requests] - text = "\n".join(lines) - if filename and filename.strip(): - with open(filename.strip(), "w", encoding="utf-8") as f: - f.write(text) - return _tool_response( - json.dumps( - { - "ok": True, - "message": f"Network requests saved to {filename}", - "filename": filename.strip(), - }, - ensure_ascii=False, - indent=2, - ), - ) - return _tool_response( - json.dumps( - {"ok": True, "requests": requests, "text": text}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_run_code(page_id: str, code: str) -> ToolResponse: - """Run JS in page (like eval). Use evaluate for element (ref).""" - code = (code or "").strip() - if not code: - return _tool_response( - json.dumps( - {"ok": False, "error": "code required for run_code"}, - ensure_ascii=False, - indent=2, - ), - ) - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if code.strip().startswith("(") or code.strip().startswith("function"): - if _USE_SYNC_PLAYWRIGHT: - result = await _run_sync(page.evaluate, code) - else: - result = await page.evaluate(code) - else: - if _USE_SYNC_PLAYWRIGHT: - result = await _run_sync( - page.evaluate, - f"() => {{ return ({code}); }}", - ) - else: - result = await page.evaluate(f"() => {{ return ({code}); }}") - try: - out = json.dumps( - {"ok": True, "result": result}, - ensure_ascii=False, - indent=2, - ) - except TypeError: - out = json.dumps( - {"ok": True, "result": str(result)}, - ensure_ascii=False, - indent=2, - ) - return _tool_response(out) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Run code failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_drag( - page_id: str, - start_ref: str, - end_ref: str, - start_selector: str = "", - end_selector: str = "", - start_element: str = "", # pylint: disable=unused-argument - end_element: str = "", # pylint: disable=unused-argument - frame_selector: str = "", -) -> ToolResponse: - start_ref = (start_ref or "").strip() - end_ref = (end_ref or "").strip() - start_selector = (start_selector or "").strip() - end_selector = (end_selector or "").strip() - use_refs = bool(start_ref and end_ref) - use_selectors = bool(start_selector and end_selector) - if not use_refs and not use_selectors: - return _tool_response( - json.dumps( - { - "ok": False, - "error": ("drag needs (start_ref,end_ref) or (start_sel,end_sel)"), - }, - ensure_ascii=False, - indent=2, - ), - ) - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - root = _get_root(page, page_id, frame_selector) - if use_refs: - start_locator = _get_locator_by_ref( - page, - page_id, - start_ref, - frame_selector, - ) - end_locator = _get_locator_by_ref( - page, - page_id, - end_ref, - frame_selector, - ) - if start_locator is None or end_locator is None: - return _tool_response( - json.dumps( - {"ok": False, "error": "Unknown ref for drag"}, - ensure_ascii=False, - indent=2, - ), - ) - else: - start_locator = root.locator(start_selector).first - end_locator = root.locator(end_selector).first - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(start_locator.drag_to, end_locator) - else: - await start_locator.drag_to(end_locator) - return _tool_response( - json.dumps( - {"ok": True, "message": "Drag completed"}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Drag failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_hover( - page_id: str, - ref: str = "", - element: str = "", # pylint: disable=unused-argument - selector: str = "", - frame_selector: str = "", -) -> ToolResponse: - ref = (ref or "").strip() - selector = (selector or "").strip() - if not ref and not selector: - return _tool_response( - json.dumps( - {"ok": False, "error": "hover requires ref or selector"}, - ensure_ascii=False, - indent=2, - ), - ) - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if ref: - locator = _get_locator_by_ref(page, page_id, ref, frame_selector) - if locator is None: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Unknown ref: {ref}"}, - ensure_ascii=False, - indent=2, - ), - ) - else: - root = _get_root(page, page_id, frame_selector) - locator = root.locator(selector).first - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(locator.hover) - else: - await locator.hover() - return _tool_response( - json.dumps( - {"ok": True, "message": f"Hovered {ref or selector}"}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Hover failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_select_option( - page_id: str, - ref: str = "", - element: str = "", # pylint: disable=unused-argument - values_json: str = "", - frame_selector: str = "", -) -> ToolResponse: - ref = (ref or "").strip() - values = _parse_json_param(values_json, []) - if not isinstance(values, list): - values = [values] if values is not None else [] - if not ref: - return _tool_response( - json.dumps( - {"ok": False, "error": "ref required for select_option"}, - ensure_ascii=False, - indent=2, - ), - ) - if not values: - return _tool_response( - json.dumps( - { - "ok": False, - "error": "values required (JSON array or comma-separated)", - }, - ensure_ascii=False, - indent=2, - ), - ) - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - locator = _get_locator_by_ref(page, page_id, ref, frame_selector) - if locator is None: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Unknown ref: {ref}"}, - ensure_ascii=False, - indent=2, - ), - ) - if _USE_SYNC_PLAYWRIGHT: - await _run_sync(locator.select_option, value=values) - else: - await locator.select_option(value=values) - return _tool_response( - json.dumps( - {"ok": True, "message": f"Selected {values}"}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Select option failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_tabs( # pylint: disable=too-many-return-statements - page_id: str, - tab_action: str, - index: int, -) -> ToolResponse: - tab_action = (tab_action or "").strip().lower() - if not tab_action: - return _tool_response( - json.dumps( - { - "ok": False, - "error": "tab_action required (list, new, close, select)", - }, - ensure_ascii=False, - indent=2, - ), - ) - pages = _state["pages"] - page_ids = list(pages.keys()) - if tab_action == "list": - return _tool_response( - json.dumps( - {"ok": True, "tabs": page_ids, "count": len(page_ids)}, - ensure_ascii=False, - indent=2, - ), - ) - if tab_action == "new": - if _USE_SYNC_PLAYWRIGHT: - if not _state["_sync_context"]: - ok = await _ensure_browser() - if not ok: - err = _state.get("_last_browser_error") or "Browser not started" - return _tool_response( - json.dumps( - {"ok": False, "error": err}, - ensure_ascii=False, - indent=2, - ), - ) - else: - if not _state["context"]: - ok = await _ensure_browser() - if not ok: - err = _state.get("_last_browser_error") or "Browser not started" - return _tool_response( - json.dumps( - {"ok": False, "error": err}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if _USE_SYNC_PLAYWRIGHT: - page = await _run_sync(_state["_sync_context"].new_page) - else: - page = await _state["context"].new_page() - new_id = _next_page_id() - _state["refs"][new_id] = {} - _state["console_logs"][new_id] = [] - _state["network_requests"][new_id] = [] - _state["pending_dialogs"][new_id] = [] - _attach_page_listeners(page, new_id) - _state["pages"][new_id] = page - _state["current_page_id"] = new_id - return _tool_response( - json.dumps( - { - "ok": True, - "page_id": new_id, - "tabs": list(_state["pages"].keys()), - }, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"New tab failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) - if tab_action == "close": - target_id = page_ids[index] if 0 <= index < len(page_ids) else page_id - return await _action_close(target_id) - if tab_action == "select": - target_id = page_ids[index] if 0 <= index < len(page_ids) else page_id - _state["current_page_id"] = target_id - return _tool_response( - json.dumps( - { - "ok": True, - "message": f"Use page_id={target_id} for later actions", - "page_id": target_id, - }, - ensure_ascii=False, - indent=2, - ), - ) - return _tool_response( - json.dumps( - {"ok": False, "error": f"Unknown tab_action: {tab_action}"}, - ensure_ascii=False, - indent=2, - ), - ) - - -async def _action_wait_for( - page_id: str, - wait_time: float, - text: str, - text_gone: str, -) -> ToolResponse: - page = _get_page(page_id) - if not page: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Page '{page_id}' not found"}, - ensure_ascii=False, - indent=2, - ), - ) - try: - if wait_time and wait_time > 0: - await asyncio.sleep(wait_time) - text = (text or "").strip() - text_gone = (text_gone or "").strip() - if text: - locator = page.get_by_text(text) - if _USE_SYNC_PLAYWRIGHT: - await _run_sync( - locator.wait_for, - state="visible", - timeout=30000, - ) - else: - await locator.wait_for( - state="visible", - timeout=30000, - ) - if text_gone: - locator = page.get_by_text(text_gone) - if _USE_SYNC_PLAYWRIGHT: - await _run_sync( - locator.wait_for, - state="hidden", - timeout=30000, - ) - else: - await locator.wait_for( - state="hidden", - timeout=30000, - ) - return _tool_response( - json.dumps( - {"ok": True, "message": "Wait completed"}, - ensure_ascii=False, - indent=2, - ), - ) - except Exception as e: - return _tool_response( - json.dumps( - {"ok": False, "error": f"Wait failed: {e!s}"}, - ensure_ascii=False, - indent=2, - ), - ) diff --git a/reme/memory/file_based/tools/browser_snapshot.py b/reme/memory/file_based/tools/browser_snapshot.py deleted file mode 100644 index 11ab8885..00000000 --- a/reme/memory/file_based/tools/browser_snapshot.py +++ /dev/null @@ -1,248 +0,0 @@ -# -*- coding: utf-8 -*- -"""Build role snapshot + refs from Playwright aria_snapshot.""" - -import re -from typing import Any - -INTERACTIVE_ROLES = frozenset( - { - "button", - "link", - "textbox", - "checkbox", - "radio", - "combobox", - "listbox", - "menuitem", - "menuitemcheckbox", - "menuitemradio", - "option", - "searchbox", - "slider", - "spinbutton", - "switch", - "tab", - "treeitem", - }, -) - -CONTENT_ROLES = frozenset( - { - "heading", - "cell", - "gridcell", - "columnheader", - "rowheader", - "listitem", - "article", - "region", - "main", - "navigation", - }, -) - -STRUCTURAL_ROLES = frozenset( - { - "generic", - "group", - "list", - "table", - "row", - "rowgroup", - "grid", - "treegrid", - "menu", - "menubar", - "toolbar", - "tablist", - "tree", - "directory", - "document", - "application", - "presentation", - "none", - }, -) - - -def _get_indent_level(line: str) -> int: - m = re.match(r"^(\s*)", line) - return int(len(m.group(1)) / 2) if m else 0 - - -def _create_tracker() -> dict[str, Any]: - counts: dict[str, int] = {} - refs_by_key: dict[str, list[str]] = {} - - def get_key(role: str, name: str | None) -> str: - return f"{role}:{name or ''}" - - def get_next_index(role: str, name: str | None) -> int: - key = get_key(role, name) - current = counts.get(key, 0) - counts[key] = current + 1 - return current - - def track_ref(role: str, name: str | None, ref: str) -> None: - key = get_key(role, name) - refs_by_key.setdefault(key, []).append(ref) - - def get_duplicate_keys() -> set[str]: - return {k for k, refs in refs_by_key.items() if len(refs) > 1} - - return { - "get_next_index": get_next_index, - "track_ref": track_ref, - "get_duplicate_keys": get_duplicate_keys, - "get_key": get_key, - } - - -def _remove_nth_from_non_duplicates( - refs: dict[str, dict], - tracker: dict, -) -> None: - dup_keys = tracker["get_duplicate_keys"]() - for _, data in list(refs.items()): - key = tracker["get_key"](data["role"], data.get("name")) - if key not in dup_keys and "nth" in data: - del data["nth"] - - -def _compact_tree(tree: str) -> str: - lines = tree.split("\n") - result = [] - for i, line in enumerate(lines): - if "[ref=" in line: - result.append(line) - continue - if ":" in line and not line.rstrip().endswith(":"): - result.append(line) - continue - current_indent = _get_indent_level(line) - has_relevant = False - for j in range(i + 1, len(lines)): - if _get_indent_level(lines[j]) <= current_indent: - break - if "[ref=" in lines[j]: - has_relevant = True - break - if has_relevant: - result.append(line) - return "\n".join(result) - - -def _process_line( # pylint: disable=too-many-return-statements - line: str, - refs: dict[str, dict], - options: dict[str, Any], - tracker: dict, - next_ref: Any, -) -> str | None: - depth = _get_indent_level(line) - max_depth_val = options.get("maxDepth") - if max_depth_val is not None and depth > max_depth_val: - return None - - m = re.match(r'^(\s*-\s*)(\w+)(?:\s+"([^"]*)")?(.*)$', line) - if not m: - return None if options.get("interactive") else line - - prefix, role_raw, name, suffix = m.groups() - if role_raw.startswith("/"): - return None if options.get("interactive") else line - - role = role_raw.lower() - is_interactive = role in INTERACTIVE_ROLES - is_content = role in CONTENT_ROLES - is_structural = role in STRUCTURAL_ROLES - - if options.get("interactive") and not is_interactive: - return None - if options.get("compact") and is_structural and not name: - return None - - should_have_ref = is_interactive or (is_content and name) - if not should_have_ref: - return line - - ref = next_ref() - nth = tracker["get_next_index"](role, name) - tracker["track_ref"](role, name, ref) - refs[ref] = {"role": role, "name": name, "nth": nth} - - enhanced = f"{prefix}{role_raw}" - if name: - enhanced += f' "{name}"' - enhanced += f" [ref={ref}]" - if nth is not None and nth > 0: - enhanced += f" [nth={nth}]" - if suffix: - enhanced += suffix - return enhanced - - -def build_role_snapshot_from_aria( - aria_snapshot: str, - *, - interactive: bool = False, - compact: bool = False, - max_depth: int | None = None, -) -> tuple[str, dict[str, dict]]: - """Build snapshot + refs from Playwright locator.aria_snapshot() output.""" - options = { - "interactive": interactive, - "compact": compact, - "maxDepth": max_depth, - } - lines = aria_snapshot.split("\n") - refs: dict[str, dict] = {} - tracker = _create_tracker() - counter = [0] - - def next_ref() -> str: - counter[0] += 1 - return f"e{counter[0]}" - - if options.get("interactive"): - result_lines = [] - for line in lines: - depth = _get_indent_level(line) - max_d = options.get("maxDepth") - if max_d is not None and depth > max_d: - continue - m = re.match(r'^(\s*-\s*)(\w+)(?:\s+"([^"]*)")?(.*)$', line) - if not m: - continue - _, role_raw, name, suffix = m.groups() - if role_raw.startswith("/"): - continue - role = role_raw.lower() - if role not in INTERACTIVE_ROLES: - continue - ref = next_ref() - nth = tracker["get_next_index"](role, name) - tracker["track_ref"](role, name, ref) - refs[ref] = {"role": role, "name": name, "nth": nth} - enhanced = f"- {role_raw}" - if name: - enhanced += f' "{name}"' - enhanced += f" [ref={ref}]" - if nth is not None and nth > 0: - enhanced += f" [nth={nth}]" - if "[" in suffix: - enhanced += suffix - result_lines.append(enhanced) - _remove_nth_from_non_duplicates(refs, tracker) - snapshot = "\n".join(result_lines) or "(no interactive elements)" - return snapshot, refs - - result_lines = [] - for line in lines: - processed = _process_line(line, refs, options, tracker, next_ref) - if processed is not None: - result_lines.append(processed) - _remove_nth_from_non_duplicates(refs, tracker) - tree = "\n".join(result_lines) or "(empty)" - snapshot = _compact_tree(tree) if options.get("compact") else tree - return snapshot, refs diff --git a/reme/memory/file_based/tools/file_io.py b/reme/memory/file_based/tools/file_io.py deleted file mode 100644 index 3e28addc..00000000 --- a/reme/memory/file_based/tools/file_io.py +++ /dev/null @@ -1,360 +0,0 @@ -"""File I/O operations with a configurable working directory.""" - -import os -from pathlib import Path -from typing import Optional - -from agentscope.message import TextBlock -from agentscope.tool import ToolResponse - -from ..utils import read_file_safe, truncate_text_output, TRUNCATION_NOTICE_MARKER - - -class FileIO: - """File I/O operations with a configurable working directory.""" - - def __init__(self, working_dir: str | Path): - """Initialize FileIO with a working directory. - - Args: - working_dir (`str`): - The working directory for resolving relative paths. - """ - self.working_dir = Path(working_dir) - - def _resolve_file_path(self, file_path: str) -> str: - """Resolve file path: use absolute path as-is, - resolve relative path from working_dir. - - Args: - file_path: The input file path (absolute or relative). - - Returns: - The resolved absolute file path as string. - """ - path = Path(file_path).expanduser() - if path.is_absolute(): - return str(path) - else: - return str(self.working_dir / file_path) - - async def read_file( # pylint: disable=too-many-return-statements - self, - file_path: str, - start_line: Optional[int] = None, - end_line: Optional[int] = None, - ) -> ToolResponse: - """Read a file. Relative paths resolve from WORKING_DIR. - - Use start_line/end_line to read a specific line range (output includes - line numbers). Omit both to read the full file. - - Args: - file_path (`str`): - Path to the file. - start_line (`int`, optional): - First line to read (1-based, inclusive). - end_line (`int`, optional): - Last line to read (1-based, inclusive). - """ - - # Convert start_line/end_line to int if they are strings - if start_line is not None: - try: - start_line = int(start_line) - except (ValueError, TypeError): - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: start_line must be an integer, got {start_line!r}.", - ), - ], - ) - - if end_line is not None: - try: - end_line = int(end_line) - except (ValueError, TypeError): - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: end_line must be an integer, got {end_line!r}.", - ), - ], - ) - - file_path = self._resolve_file_path(file_path) - - if not os.path.exists(file_path): - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: The file {file_path} does not exist.", - ), - ], - ) - - if not os.path.isfile(file_path): - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: The path {file_path} is not a file.", - ), - ], - ) - - try: - content = read_file_safe(file_path) - all_lines = content.split("\n") - total = len(all_lines) - - # Determine read range - s = max(1, start_line if start_line is not None else 1) - e = min(total, end_line if end_line is not None else total) - - if s > total: - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: start_line {s} exceeds file length ({total} lines).", - ), - ], - ) - - if s > e: - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: start_line ({s}) > end_line ({e}).", - ), - ], - ) - - # Extract selected lines - selected_content = "\n".join(all_lines[s - 1 : e]) - - # Apply smart truncation (consistent with shell output format) - text = truncate_text_output( - selected_content, - start_line=s, - total_lines=total, - file_path=file_path, - ) - - # Add continuation hint if partial read without truncation. - # Use TRUNCATION_NOTICE_MARKER format so ToolResultCompactor can - # re-truncate with the correct start_line when compacting old messages. - if text == selected_content and e < total: - content_bytes = len(text.encode("utf-8")) - notice = ( - TRUNCATION_NOTICE_MARKER + f"\nThe output above was truncated." - f"\nThe full content is saved to the file " - f"and contains {total} lines in total." - f"\nThis excerpt starts at line {s} and " - f"covers the next {content_bytes} bytes." - "\nIf the current content is not enough, " - f"call `read_file` with file_path={file_path} start_line={e + 1} to read more." - ) - text = text + notice - - return ToolResponse( - content=[TextBlock(type="text", text=text)], - ) - - except Exception as e: - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: Read file failed due to \n{e}", - ), - ], - ) - - async def write_file( - self, - file_path: str, - content: str, - ) -> ToolResponse: - """Create or overwrite a file. Relative paths resolve from working_dir. - - Args: - file_path (`str`): - Path to the file. - content (`str`): - Content to write. - """ - if not file_path: - return ToolResponse( - content=[ - TextBlock( - type="text", - text="Error: No `file_path` provided.", - ), - ], - ) - - file_path = self._resolve_file_path(file_path) - - try: - with open(file_path, "w", encoding="utf-8") as file: - file.write(content) - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Wrote {len(content)} bytes to {file_path}.", - ), - ], - ) - except Exception as e: - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: Write file failed due to \n{e}", - ), - ], - ) - - # pylint: disable=too-many-return-statements - async def edit_file( - self, - file_path: str, - old_text: str, - new_text: str, - ) -> ToolResponse: - """Find-and-replace text in a file. All occurrences of old_text are - replaced with new_text. Relative paths resolve from working_dir. - - Args: - file_path (`str`): - Path to the file. - old_text (`str`): - Exact text to find. - new_text (`str`): - Replacement text. - """ - if not file_path: - return ToolResponse( - content=[ - TextBlock( - type="text", - text="Error: No `file_path` provided.", - ), - ], - ) - - resolved_path = self._resolve_file_path(file_path) - - if not os.path.exists(resolved_path): - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: The file {resolved_path} does not exist.", - ), - ], - ) - - if not os.path.isfile(resolved_path): - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: The path {resolved_path} is not a file.", - ), - ], - ) - - try: - content = read_file_safe(resolved_path) - except Exception as e: - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: Read file failed due to \n{e}", - ), - ], - ) - - if old_text not in content: - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: The text to replace was not found in {file_path}.", - ), - ], - ) - - new_content = content.replace(old_text, new_text) - write_response = await self.write_file(file_path=resolved_path, content=new_content) - - if write_response.content and len(write_response.content) > 0: - write_text = write_response.content[0].get("text", "") - if write_text.startswith("Error:"): - return write_response - - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Successfully replaced text in {file_path}.", - ), - ], - ) - - async def append_file( - self, - file_path: str, - content: str, - ) -> ToolResponse: - """Append content to the end of a file. Relative paths resolve from - working_dir. - - Args: - file_path (`str`): - Path to the file. - content (`str`): - Content to append. - """ - if not file_path: - return ToolResponse( - content=[ - TextBlock( - type="text", - text="Error: No `file_path` provided.", - ), - ], - ) - - file_path = self._resolve_file_path(file_path) - - try: - with open(file_path, "a", encoding="utf-8") as file: - file.write(content) - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Appended {len(content)} bytes to {file_path}.", - ), - ], - ) - except Exception as e: - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: Append file failed due to \n{e}", - ), - ], - ) diff --git a/reme/memory/file_based/tools/memory_get.py b/reme/memory/file_based/tools/memory_get.py deleted file mode 100644 index 9cc16968..00000000 --- a/reme/memory/file_based/tools/memory_get.py +++ /dev/null @@ -1,116 +0,0 @@ -"""Memory get tool for reading specific snippets from memory files.""" - -import os -from pathlib import Path - - -from ....core import RuntimeContext -from ....core.op import BaseTool -from ....core.schema import ToolCall - -from ....core.utils import get_logger - -logger = get_logger() - - -class MemoryGet(BaseTool): - """Read specific snippets from memory files.""" - - def __init__(self, cwd: str | None = None, **kwargs): - """Initialize memory get tool.""" - kwargs.setdefault("max_retries", 1) - kwargs.setdefault("raise_exception", False) - super().__init__(**kwargs) - self.cwd = cwd or os.getcwd() - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": ( - "Safe snippet read from MEMORY.md, memory/*.md with optional offset/limit; " - "use after memory_search to pull only the needed lines and keep context small." - ), - "parameters": { - "type": "object", - "properties": { - "path": { - "type": "string", - "description": "Path to the memory file to read (relative or absolute)", - }, - "offset": { - "type": "integer", - "description": "Starting line number (1-indexed, optional)", - }, - "limit": { - "type": "integer", - "description": "Number of lines to read from the starting line (optional)", - }, - }, - "required": ["path"], - }, - }, - ) - - async def execute(self) -> str: - """Execute the memory get operation.""" - raw_path: str = self.context.path.strip() - offset: int | None = self.context.get("offset", None) - limit: int | None = self.context.get("limit", None) - - if os.path.isabs(raw_path): - abs_path = os.path.abspath(raw_path) - else: - abs_path = os.path.abspath(os.path.join(self.cwd, raw_path)) - assert abs_path.lower().endswith(".md") - - # Check file exists, is not a symlink, and is a regular file - file_path = Path(abs_path) - assert ( - file_path.exists() and not file_path.is_symlink() and file_path.is_file() - ), f"File not found or not a regular file: {abs_path}" - - with open(abs_path, "r", encoding="utf-8") as f: - content = f.read() - - if offset is None and limit is None: - return content - - lines = content.split("\n") - total_lines = len(lines) - - # Validate and normalize offset (1-indexed) - start = offset if offset is not None else 1 - assert start >= 1, f"offset must be >= 1, got {start}" - assert start <= total_lines, f"offset {start} exceeds total lines {total_lines}" - - # Validate and calculate count - if limit is not None: - assert limit > 0, f"limit must be positive, got {limit}" - count = limit - else: - # Read from start to end of file - count = total_lines - start + 1 - - # Extract slice (1-indexed to 0-indexed conversion) - selected = lines[start - 1 : start - 1 + count] - return "\n".join(selected) - - async def call(self, context: RuntimeContext = None, **kwargs): - """Execute the tool with unified error handling. - - This method catches all exceptions and returns error messages - to the LLM instead of raising them. - """ - self.context = RuntimeContext.from_context(context, **kwargs) - - try: - await self.before_execute() - response = await self.execute() - response = await self.after_execute(response) - return response - - except Exception as e: - # Return error message to LLM instead of raising - error_msg = f"{self.__class__.__name__} failed: {str(e)}" - logger.exception(error_msg) - return await self.after_execute(error_msg) diff --git a/reme/memory/file_based/tools/memory_search.py b/reme/memory/file_based/tools/memory_search.py deleted file mode 100644 index 927a6fc1..00000000 --- a/reme/memory/file_based/tools/memory_search.py +++ /dev/null @@ -1,114 +0,0 @@ -"""Memory search tool for semantic search in memory files.""" - -import json - - -from ....core.enumeration import MemorySource -from ....core.op import BaseTool -from ....core.runtime_context import RuntimeContext -from ....core.schema import ToolCall -from ....core.utils import get_logger - -logger = get_logger() - - -class MemorySearch(BaseTool): - """Semantically search MEMORY.md and memory files.""" - - def __init__( - self, - sources: list[MemorySource] | None = None, - min_score: float = 0.1, - max_results: int = 5, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - **kwargs, - ): - """Initialize memory search tool.""" - assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}" - kwargs.setdefault("max_retries", 1) - kwargs.setdefault("raise_exception", False) - super().__init__(**kwargs) - self.sources = sources or [MemorySource.MEMORY] - self.min_score = min_score - self.max_results = max_results - self.vector_weight = vector_weight - self.candidate_multiplier = candidate_multiplier - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": ( - "Mandatory recall step: semantically search MEMORY.md + memory/*.md " - "(and optional session transcripts) before answering questions about " - "prior work, decisions, dates, people, preferences, or todos; returns " - "top snippets with path + lines." - ), - "parameters": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "The semantic search query to find relevant memory snippets", - }, - "max_results": { - "type": "integer", - "description": "Maximum number of search results to return (optional), default 5", - }, - "min_score": { - "type": "number", - "description": "Minimum similarity score threshold for results (optional), default 0.1", - }, - }, - "required": ["query"], - }, - }, - ) - - async def execute(self) -> str: - """Execute the memory search operation.""" - query: str = self.context.query.strip() - min_score: float = self.context.get("min_score", self.min_score) - max_results: int = self.context.get("max_results", self.max_results) - - assert query, "Query cannot be empty" - assert ( - isinstance(min_score, float) and 0.0 <= min_score <= 1.0 - ), f"min_score must be between 0 and 1, got {min_score}" - assert ( - isinstance(max_results, int) and max_results > 0 - ), f"max_results must be a positive integer, got {max_results}" - - # Use hybrid_search from file_store - results = await self.file_store.hybrid_search( - query=query, - limit=max_results, - sources=self.sources, - vector_weight=self.vector_weight, - candidate_multiplier=self.candidate_multiplier, - ) - - # Filter by min_score - results = [r for r in results if r.score >= min_score] - - return json.dumps([result.model_dump(exclude_none=True) for result in results], indent=2, ensure_ascii=False) - - async def call(self, context: RuntimeContext = None, **kwargs): - """Execute the tool with unified error handling. - - This method catches all exceptions and returns error messages - to the LLM instead of raising them. - """ - self.context = RuntimeContext.from_context(context, **kwargs) - - try: - await self.before_execute() - response = await self.execute() - response = await self.after_execute(response) - return response - - except Exception as e: - # Return error message to LLM instead of raising - error_msg = f"{self.__class__.__name__} failed: {str(e)}" - logger.exception(error_msg) - return await self.after_execute(error_msg) diff --git a/reme/memory/file_based/tools/shell.py b/reme/memory/file_based/tools/shell.py deleted file mode 100644 index cce138f9..00000000 --- a/reme/memory/file_based/tools/shell.py +++ /dev/null @@ -1,223 +0,0 @@ -# -*- coding: utf-8 -*- -# flake8: noqa: E501 -# pylint: disable=line-too-long -"""The shell command tool.""" - -import asyncio -import locale -import subprocess -import sys -from pathlib import Path - -from agentscope.message import TextBlock -from agentscope.tool import ToolResponse - - -def _execute_subprocess_sync( - cmd: str, - cwd: str, - timeout: int, -) -> tuple[int, str, str]: - """Execute subprocess synchronously in a thread. - - This function runs in a separate thread to avoid Windows asyncio - subprocess limitations. - - Args: - cmd (`str`): - The shell command to execute. - cwd (`str`): - The working directory for the command execution. - timeout (`int`): - The maximum time (in seconds) allowed for the command to run. - - Returns: - `tuple[int, str, str]`: - A tuple containing the return code, standard output, and - standard error of the executed command. If timeout occurs, the - return code will be -1 and stderr will contain timeout information. - """ - try: - result = subprocess.run( - cmd, - shell=True, - capture_output=True, - text=True, - cwd=cwd, - timeout=timeout, - encoding=locale.getpreferredencoding(False) or "utf-8", - errors="replace", - check=True, - ) - return ( - result.returncode, - result.stdout.strip("\n"), - result.stderr.strip("\n"), - ) - except subprocess.TimeoutExpired: - return ( - -1, - "", - f"Command execution exceeded the timeout of {timeout} seconds.", - ) - except Exception as e: - return -1, "", str(e) - - -class Shell: - """Shell command execution with a configurable working directory.""" - - def __init__(self, working_dir: str | Path): - """Initialize Shell with a working directory. - - Args: - working_dir (`str | Path`): - The working directory for command execution. - """ - self.working_dir = Path(working_dir) - - # pylint: disable=too-many-branches, too-many-statements - async def execute_shell_command( - self, - command: str, - timeout: int = 60, - ) -> ToolResponse: - """Execute given command and return the return code, standard output and - error within , and - tags. - - Args: - command (`str`): - The shell command to execute. - timeout (`int`, defaults to `60`): - The maximum time (in seconds) allowed for the command to run. - Default is 60 seconds. - - Returns: - `ToolResponse`: - The tool response containing the return code, standard output, and - standard error of the executed command. If timeout occurs, the - return code will be -1 and stderr will contain timeout information. - """ - - cmd = (command or "").strip() - - # Set working directory - working_dir = self.working_dir - - try: - if sys.platform == "win32": - # Windows: use thread pool to avoid asyncio subprocess limitations - returncode, stdout_str, stderr_str = await asyncio.to_thread( - _execute_subprocess_sync, - cmd, - str(working_dir), - timeout, - ) - else: - proc = await asyncio.create_subprocess_shell( - cmd, - stdout=asyncio.subprocess.PIPE, - stderr=asyncio.subprocess.PIPE, - bufsize=0, - cwd=str(working_dir), - ) - - try: - # Apply timeout to communicate directly; wait()+communicate() - # can hang if descendants keep stdout/stderr pipes open. - stdout, stderr = await asyncio.wait_for( - proc.communicate(), - timeout=timeout, - ) - encoding = locale.getpreferredencoding(False) or "utf-8" - stdout_str = stdout.decode(encoding, errors="replace").strip( - "\n", - ) - stderr_str = stderr.decode(encoding, errors="replace").strip( - "\n", - ) - returncode = proc.returncode - - except asyncio.TimeoutError: - # Handle timeout - stderr_suffix = ( - f"⚠️ TimeoutError: The command execution exceeded " - f"the timeout of {timeout} seconds. " - f"Please consider increasing the timeout value if this command " - f"requires more time to complete." - ) - returncode = -1 - try: - proc.terminate() - # Wait a bit for graceful termination - try: - await asyncio.wait_for(proc.wait(), timeout=1) - except asyncio.TimeoutError: - # Force kill if graceful termination fails - proc.kill() - await proc.wait() - - # Avoid hanging forever while draining pipes after timeout. - try: - stdout, stderr = await asyncio.wait_for( - proc.communicate(), - timeout=1, - ) - except asyncio.TimeoutError: - stdout, stderr = b"", b"" - encoding = locale.getpreferredencoding(False) or "utf-8" - stdout_str = stdout.decode( - encoding, - errors="replace", - ).strip( - "\n", - ) - stderr_str = stderr.decode( - encoding, - errors="replace", - ).strip( - "\n", - ) - if stderr_str: - stderr_str += f"\n{stderr_suffix}" - else: - stderr_str = stderr_suffix - except ProcessLookupError: - stdout_str = "" - stderr_str = stderr_suffix - - # Format the response in a human-friendly way - if returncode == 0: - # Success case: just show the output - if stdout_str: - response_text = stdout_str - else: - response_text = "Command executed successfully (no output)." - else: - # Error case: show detailed information - response_parts = [f"Command failed with exit code {returncode}."] - if stdout_str: - response_parts.append(f"\n[stdout]\n{stdout_str}") - if stderr_str: - response_parts.append(f"\n[stderr]\n{stderr_str}") - response_text = "".join(response_parts) - - return ToolResponse( - content=[ - TextBlock( - type="text", - text=response_text, - ), - ], - ) - - except Exception as e: - return ToolResponse( - content=[ - TextBlock( - type="text", - text=f"Error: Shell command execution failed due to \n{e}", - ), - ], - ) diff --git a/reme/memory/file_based/utils/__init__.py b/reme/memory/file_based/utils/__init__.py deleted file mode 100644 index 769a8a66..00000000 --- a/reme/memory/file_based/utils/__init__.py +++ /dev/null @@ -1,12 +0,0 @@ -"""utils""" - -from .as_msg_handler import AsMsgHandler -from .file_utils import truncate_text_output, read_file_safe, DEFAULT_MAX_BYTES, TRUNCATION_NOTICE_MARKER - -__all__ = [ - "AsMsgHandler", - "truncate_text_output", - "read_file_safe", - "DEFAULT_MAX_BYTES", - "TRUNCATION_NOTICE_MARKER", -] diff --git a/reme/memory/file_based/utils/as_msg_handler.py b/reme/memory/file_based/utils/as_msg_handler.py deleted file mode 100644 index e175a402..00000000 --- a/reme/memory/file_based/utils/as_msg_handler.py +++ /dev/null @@ -1,398 +0,0 @@ -"""Handler for AgentScope message processing, token counting, and context management.""" - -import json - -from agentscope.message import Msg -from agentscope.token import HuggingFaceTokenCounter - -from ....core.schema import AsMsgStat, AsBlockStat -from ....core.utils import get_logger - -logger = get_logger() - - -class AsMsgHandler: - """Handles token counting, formatting, and context compaction for AgentScope messages.""" - - def __init__(self, token_counter: HuggingFaceTokenCounter): - self._token_counter = token_counter - - async def count_str_token(self, text: str) -> int: - """Count tokens in a string. - - Args: - text: The text to count tokens for. - - Returns: - The number of tokens in the text. - """ - if not text: - return 0 - - try: - token_count = await self._token_counter.count(messages=[], text=text) - assert token_count > 0, "Invalid token count" - return token_count - - except Exception as e: - estimated_tokens = int(len(text.encode("utf-8")) / 3.75) - logger.warning(f"Failed to count string tokens: {text}, e={e}") - return estimated_tokens - - async def _format_tool_result_output(self, output: str | list[dict]) -> tuple[str, int]: - """Convert tool result output to string.""" - if isinstance(output, str): - return output, await self.count_str_token(output) - - textual_parts = [] - total_token_count = 0 - for block in output: - try: - if not isinstance(block, dict) or "type" not in block: - logger.warning( - f"Invalid block: {block}, expected a dict with 'type' key, skipped.", - ) - continue - - block_type = block["type"] - - if block_type == "text": - textual_parts.append(block.get("text", "")) - total_token_count += await self.count_str_token(textual_parts[-1]) - - elif block_type in ["image", "audio", "video"]: - source = block.get("source", {}) - if source.get("type") == "base64": - data = source.get("data", "") - total_token_count += len(data) // 4 if data else 10 - else: - url = source.get("url", "") - total_token_count += await self.count_str_token(url) if url else 10 - textual_parts.append(f"[{block_type}] {url}") - - elif block_type == "file": - file_path = block.get("path", "") or block.get("url", "") - file_name = block.get("name", file_path) - textual_parts.append(f"[file] {file_name}: {file_path}") - total_token_count += await self.count_str_token(file_path) - - else: - logger.warning( - f"Unsupported block type '{block_type}' in tool result, skipped.", - ) - - except Exception as e: - logger.warning( - f"Failed to process block {block}: {e}, skipped.", - ) - - return "\n".join(textual_parts), total_token_count - - async def stat_message(self, message: Msg) -> AsMsgStat: - """Analyze a message and generate block statistics.""" - blocks = [] - if isinstance(message.content, str): - blocks.append( - AsBlockStat( - block_type="text", - text=message.content, - token_count=await self.count_str_token(message.content), - ), - ) - return AsMsgStat( - name=message.name or message.role, - role=message.role, - content=blocks, - timestamp=message.timestamp or "", - metadata=message.metadata or {}, - ) - - for block in message.content: - block_type = block.get("type", "unknown") - - if block_type == "text": - text = block.get("text", "") - token_count = await self.count_str_token(text) - blocks.append( - AsBlockStat( - block_type=block_type, - text=text, - token_count=token_count, - ), - ) - - elif block_type == "thinking": - thinking = block.get("thinking", "") - token_count = await self.count_str_token(thinking) - blocks.append( - AsBlockStat( - block_type=block_type, - text=thinking, - token_count=token_count, - ), - ) - - elif block_type in ("image", "audio", "video"): - source = block.get("source", {}) - url = source.get("url", "") - if source.get("type") == "base64": - data = source.get("data", "") - token_count = len(data) // 4 if data else 10 - else: - token_count = await self.count_str_token(url) if url else 10 - blocks.append( - AsBlockStat( - block_type=block_type, - text="", - token_count=token_count, - media_url=url, - ), - ) - - elif block_type == "tool_use": - tool_name = block.get("name", "") - tool_input = block.get("input", "") - try: - input_str = json.dumps(tool_input, ensure_ascii=False) - except (TypeError, ValueError): - input_str = str(tool_input) - token_count = await self.count_str_token(tool_name + input_str) - blocks.append( - AsBlockStat( - block_type=block_type, - text="", - token_count=token_count, - tool_name=tool_name, - tool_input=input_str, - ), - ) - - elif block_type == "tool_result": - tool_name = block.get("name", "") - output = block.get("output", "") - formatted_output, token_count = await self._format_tool_result_output(output) - blocks.append( - AsBlockStat( - block_type=block_type, - text="", - token_count=token_count, - tool_name=tool_name, - tool_output=formatted_output, - ), - ) - - else: - logger.warning(f"Unsupported block type {block_type}, skipped.") - - return AsMsgStat( - name=message.name or message.role, - role=message.role, - content=blocks, - timestamp=message.timestamp or "", - metadata=message.metadata or {}, - ) - - async def count_msgs_token(self, messages: list[Msg]) -> int: - """Count total token count of a list of messages.""" - total = 0 - for msg in messages: - stat = await self.stat_message(msg) - total += stat.total_tokens - return total - - async def format_msgs_to_str( - self, - messages: list[Msg], - memory_compact_threshold: int, - include_thinking: bool = True, - ) -> str: - """Format list of messages to a single formatted string. - - Messages are processed in reverse order (newest first) and older - messages are skipped when token count exceeds memory_compact_threshold. - - Args: - messages: List of Msg objects to format. - memory_compact_threshold: Maximum token count before skipping older messages. - include_thinking: Whether to include thinking blocks in output. - """ - if not messages: - return "" - - formatted_parts: list[str] = [] - total_token_count = 0 - - for i in range(len(messages) - 1, -1, -1): - stat = await self.stat_message(messages[i]) - formatted_content = stat.format(include_thinking=include_thinking) - content_token_count = await self.count_str_token(formatted_content) - - is_latest = i == len(messages) - 1 - if not is_latest and total_token_count + content_token_count > memory_compact_threshold: - logger.info( - f"Skipping older messages: adding {content_token_count} tokens would exceed threshold " - f"{memory_compact_threshold} (current: {total_token_count})", - ) - break - - if is_latest and content_token_count > memory_compact_threshold: - logger.warning( - f"Latest message alone ({content_token_count} tokens) exceeds threshold " - f"{memory_compact_threshold}, including it anyway.", - ) - - formatted_parts.append(formatted_content) - total_token_count += content_token_count - - formatted_parts.reverse() - return "\n\n".join(formatted_parts) - - @staticmethod - def validate_tool_ids_alignment(messages: list[Msg]) -> bool: - """Check if tool_use_ids and tool_result_ids are properly aligned. - - Args: - messages: List of Msg objects to validate. - - Returns: - True if all tool_use ids have corresponding tool_result ids and vice versa. - """ - tool_use_ids: set[str] = set() - tool_result_ids: set[str] = set() - - for msg in messages: - for block in msg.get_content_blocks("tool_use"): - if tool_id := block.get("id"): - tool_use_ids.add(tool_id) - for block in msg.get_content_blocks("tool_result"): - if tool_id := block.get("id"): - tool_result_ids.add(tool_id) - - return tool_use_ids == tool_result_ids - - async def context_check( - self, - messages: list[Msg], - memory_compact_threshold: int, - memory_compact_reserve: int, - ) -> tuple[list[Msg], list[Msg], bool]: - """Check if context exceeds threshold and split messages accordingly. - - Only when total tokens exceed memory_compact_threshold, messages are split into - messages_to_keep (within reserve limit) and messages_to_compact (older messages). - - Args: - messages: List of Msg objects to check. - memory_compact_threshold: Maximum token count threshold to trigger compaction. - memory_compact_reserve: Token limit for messages to keep. - - Returns: - A tuple of (messages_to_compact, messages_to_keep, tools_aligned): - - messages_to_compact: Older messages that exceed reserve limit - - messages_to_keep: Recent messages within the reserve limit - - tools_aligned: Whether tool_use and tool_result ids are aligned in messages_to_keep - """ - if not messages: - return [], [], True - - # Calculate total tokens and stats for all messages - msg_stats: list[tuple[Msg, AsMsgStat]] = [] - total_tokens = 0 - for msg in messages: - stat = await self.stat_message(msg) - msg_stats.append((msg, stat)) - total_tokens += stat.total_tokens - - # If total tokens don't exceed threshold, no split needed - if total_tokens < memory_compact_threshold: - return [], messages, True - - # Collect all tool_use ids and their message indices - # tool_use_id -> message index - tool_use_locations: dict[str, int] = {} - # tool_result_id -> message index - tool_result_locations: dict[str, int] = {} - - for idx, (msg, _) in enumerate(msg_stats): - for block in msg.get_content_blocks("tool_use"): - tool_id = block.get("id", "") - if tool_id: - tool_use_locations[tool_id] = idx - - for block in msg.get_content_blocks("tool_result"): - tool_id = block.get("id", "") - if tool_id: - tool_result_locations[tool_id] = idx - - # Iterate from the end, accumulating messages to keep within reserve limit - keep_indices: set[int] = set() - accumulated_tokens = 0 - - for i in range(len(msg_stats) - 1, -1, -1): - # Skip messages already added as tool_use dependencies to avoid double-counting tokens - if i in keep_indices: - continue - - msg, stat = msg_stats[i] - - # Check if adding this message would exceed reserve limit - if accumulated_tokens + stat.total_tokens > memory_compact_reserve: - logger.info( - f"Context check: adding message {i} with {stat.total_tokens} tokens would exceed reserve " - f"{memory_compact_reserve} (current: {accumulated_tokens})", - ) - break - - # Check tool_result dependencies - if this message has tool_result, - # we need to ensure the corresponding tool_use is also included - tool_result_ids = [ - block.get("id", "") for block in msg.get_content_blocks("tool_result") if block.get("id", "") - ] - - # Calculate extra tokens needed for dependent tool_use messages - extra_tokens = 0 - dependent_indices: set[int] = set() - - for tool_id in tool_result_ids: - if tool_id in tool_use_locations: - tool_use_idx = tool_use_locations[tool_id] - if tool_use_idx not in keep_indices and tool_use_idx != i: - dependent_indices.add(tool_use_idx) - _, dep_stat = msg_stats[tool_use_idx] - extra_tokens += dep_stat.total_tokens - - # Check if we can fit this message plus its dependencies within reserve - if accumulated_tokens + stat.total_tokens + extra_tokens > memory_compact_reserve: - logger.info( - f"Context check: message {i} requires {extra_tokens} extra tokens for tool_use dependencies, " - f"total would exceed reserve {memory_compact_reserve}", - ) - break - - # Add this message and its dependencies - keep_indices.add(i) - keep_indices.update(dependent_indices) - accumulated_tokens += stat.total_tokens + extra_tokens - - # Build final lists based on keep_indices (preserve original order) - messages_to_compact = [] - messages_to_keep = [] - - for idx, (msg, _) in enumerate(msg_stats): - if idx in keep_indices: - messages_to_keep.append(msg) - else: - messages_to_compact.append(msg) - - # Validate tool ids alignment for messages_to_keep - tools_aligned = self.validate_tool_ids_alignment(messages_to_keep) - - logger.info( - f"Context check result: {len(messages_to_compact)} messages to compact, " - f"{len(messages_to_keep)} messages to keep, " - f"total tokens: {total_tokens}, threshold: {memory_compact_threshold}, " - f"reserve: {memory_compact_reserve}, kept tokens: {accumulated_tokens}, " - f"tools_aligned: {tools_aligned}", - ) - - return messages_to_compact, messages_to_keep, tools_aligned diff --git a/reme/memory/file_based/utils/file_utils.py b/reme/memory/file_based/utils/file_utils.py deleted file mode 100644 index 50248554..00000000 --- a/reme/memory/file_based/utils/file_utils.py +++ /dev/null @@ -1,217 +0,0 @@ -# -*- coding: utf-8 -*- -"""Shared utilities for file and shell tools.""" - -import re - -from ....core.utils import get_logger - -logger = get_logger() - -# Default truncation limit -DEFAULT_MAX_BYTES = 50 * 1024 - -# Maximum file size to read into memory (1GB) -MAX_FILE_READ_BYTES = 1024 * 1024 * 1024 - -# Marker prepended to every truncation notice. -# Format: -# <<>> -# The output above was truncated. -# The full content is saved to the file and contains Z lines in total. -# This excerpt starts at line X and covers the next N bytes. -# If the current content is not enough, call `read_file` with file_path= start_line=Y to read more. -# -# Split output on this marker to recover the original (untruncated) portion: -# original = output.split(TRUNCATION_NOTICE_MARKER)[0] -TRUNCATION_NOTICE_MARKER = "<<>>" - - -def _truncate_fresh( - text: str, - start_line: int, - total_lines: int, - max_bytes: int, - file_path: str | None, - encoding: str, -) -> str: - """Truncate fresh text (no prior truncation marker) by bytes with line integrity. - - Slices at the byte boundary and appends a truncation notice with a continuation - hint so callers know which line to read next. - - Returns the original text unchanged when it fits within max_bytes, or when the - last line itself exceeds max_bytes (unhandled edge case). - """ - text_bytes = text.encode(encoding) - - # Under the byte limit — return as-is without any modification. - if len(text_bytes) <= max_bytes: - return text - - # Slice at the byte boundary. - # Assuming every single line is shorter than DEFAULT_MAX_BYTES, this cut always - # lands mid-line, guaranteeing at least one complete line before the boundary. - # Lines that exceed DEFAULT_MAX_BYTES are not handled and may be skipped entirely. - truncated = text_bytes[:max_bytes] - # Decode back to str; errors="ignore" drops any split multi-byte character - # at the cut boundary without raising an exception. - result = truncated.decode(encoding, errors="ignore") - - # Count '\n' characters to determine how many complete lines are included. - # The tail after the final '\n' is a partial line that will be covered by - # the next read starting at next_line. - newline_count = result.count("\n") - - # Compute the first line number not yet fully included in this chunk. - # max(1, ...) prevents next_line from equaling start_line when a single line - # exceeds max_bytes (newline_count == 0), which would make the caller retry - # the same range indefinitely. - next_line = start_line + max(1, newline_count) - - if next_line <= total_lines: - # Truncation fell before the last line — continue reading from next_line. - read_from = next_line - elif start_line < total_lines: - # next_line overshot total_lines, meaning the cut landed inside the last line. - # Re-read from the start of the last line so the caller gets it in full. - read_from = total_lines - else: - # start_line == total_lines: the last line itself exceeds DEFAULT_MAX_BYTES. - # This case is outside our handled range — return without a truncation notice. - return result - - notice = ( - TRUNCATION_NOTICE_MARKER + f"\nThe output above was truncated." - f"\nThe full content is saved to the file and contains {total_lines} lines in total." - f"\nThis excerpt starts at line {start_line} and covers the next {max_bytes} bytes." - f"\nIf the current content is not enough, call `read_file` with file_path={file_path or ''} " - f"start_line={read_from} to read more." - ) - - return result + notice - - -def _retruncate( - text: str, - max_bytes: int, - encoding: str, -) -> str: - """Re-truncate text that was previously truncated (contains TRUNCATION_NOTICE_MARKER). - - Extracts the original content before the marker, applies the new byte limit, and - updates the embedded notice (byte count and continuation line number) via regex. - - Returns the original text unchanged when: - - the content already fits within max_bytes (with a small slack); - - required metadata fields cannot be parsed from the existing notice. - """ - parts = text.split(TRUNCATION_NOTICE_MARKER, 1) - original_content = parts[0] - old_notice = parts[1] - - text_bytes = original_content.encode(encoding) - - # Allow a small slack to avoid unnecessary re-truncation when content is just - # barely over the limit (e.g. due to minor encoding differences). - if len(text_bytes) <= max_bytes + 100: - return text - - # Parse start_line from notice; return text unchanged if not found - start_match = re.search(r"starts at line (\d+)", old_notice) - if not start_match: - return text - start_line_parsed = int(start_match.group(1)) - - # Re-slice to the new byte limit. - # Because every line is assumed to be shorter than DEFAULT_MAX_BYTES, the cut - # always falls somewhere mid-line, so at least one complete line is preserved. - truncated_bytes = text_bytes[:max_bytes] - # errors="ignore" silently drops any incomplete multi-byte character at the cut boundary. - result = truncated_bytes.decode(encoding, errors="ignore") - # Each '\n' in result corresponds to one fully-included line; - # anything after the last '\n' is a partial line that was cut off. - newline_count = result.count("\n") - - # The next read should start at the line immediately after all complete lines. - # max(1, ...) guards against the theoretical zero-newline case - # (impossible when every line is shorter than DEFAULT_MAX_BYTES). - next_line = start_line_parsed + max(1, newline_count) - - if not re.search(r"covers the next \d+ bytes", old_notice): - return text - # _truncate_fresh always includes a continuation hint, so both fields are always present. - new_notice = re.sub(r"covers the next \d+ bytes", f"covers the next {max_bytes} bytes", old_notice) - new_notice = re.sub(r"start_line=\d+ to read more", f"start_line={next_line} to read more", new_notice) - - return result + TRUNCATION_NOTICE_MARKER + new_notice - - -def truncate_text_output( - text: str, - start_line: int = 1, - total_lines: int = 0, - max_bytes: int = DEFAULT_MAX_BYTES, - file_path: str | None = None, - encoding: str = "utf-8", -) -> str: - """Truncate file output by bytes with line integrity. - - If text is under byte limit, return as-is. - If over limit, truncate at the last complete line that fits, - allowing the next read to start from a fresh line. - - Dispatches to :func:`_truncate_fresh` for text seen for the first time, or to - :func:`_retruncate` when the text already contains a TRUNCATION_NOTICE_MARKER - from a previous pass. - - Args: - text: The output text to truncate. - start_line: The starting line number (1-based). Ignored when text already - contains a truncation notice (values are parsed from the notice instead). - total_lines: Total lines in the original file. Ignored when text already - contains a truncation notice (values are parsed from the notice instead). - max_bytes: Maximum size in bytes. - file_path: Optional file path to include in the truncation notice. - encoding: Character encoding used for byte-length calculation and decoding. - - Returns: - Truncated text with notice if truncated. - """ - if not text: - return text - if max_bytes <= 0: - return text - - try: - if TRUNCATION_NOTICE_MARKER in text: - return _retruncate(text, max_bytes=max_bytes, encoding=encoding) - else: - return _truncate_fresh( - text, - start_line=start_line, - total_lines=total_lines, - max_bytes=max_bytes, - file_path=file_path, - encoding=encoding, - ) - except Exception: - logger.warning("truncate_text_output failed, returning original text", exc_info=True) - return text - - -def read_file_safe(file_path: str, max_bytes: int = MAX_FILE_READ_BYTES) -> str: - """Read file with Unicode error handling and memory protection. - - Args: - file_path: Path to the file. - max_bytes: Maximum bytes to read into memory (default 1GB). - - Returns: - File content as string (up to max_bytes). - """ - try: - with open(file_path, "r", encoding="utf-8") as f: - return f.read(max_bytes) - except UnicodeDecodeError: - with open(file_path, "r", encoding="utf-8", errors="ignore") as f: - return f.read(max_bytes) diff --git a/reme/memory/vector_based/__init__.py b/reme/memory/vector_based/__init__.py deleted file mode 100644 index 61b45f37..00000000 --- a/reme/memory/vector_based/__init__.py +++ /dev/null @@ -1,33 +0,0 @@ -"""memory agent""" - -from .base_memory_agent import BaseMemoryAgent -from .personal.personal_retriever import PersonalRetriever -from .personal.personal_summarizer import PersonalSummarizer -from .procedural.procedural_retriever import ProceduralRetriever -from .procedural.procedural_summarizer import ProceduralSummarizer -from .reme_retriever import ReMeRetriever -from .reme_summarizer import ReMeSummarizer -from .tool_call.tool_retriever import ToolRetriever -from .tool_call.tool_summarizer import ToolSummarizer -from ...core import R - -__all__ = [ - "BaseMemoryAgent", - "PersonalRetriever", - "PersonalSummarizer", - "ProceduralRetriever", - "ProceduralSummarizer", - "ReMeRetriever", - "ReMeSummarizer", - "ToolRetriever", - "ToolSummarizer", -] - -for name in __all__: - agent_class = globals()[name] - if ( - isinstance(agent_class, type) - and issubclass(agent_class, BaseMemoryAgent) - and agent_class is not BaseMemoryAgent - ): - R.ops.register(agent_class) diff --git a/reme/memory/vector_based/base_memory_agent.py b/reme/memory/vector_based/base_memory_agent.py deleted file mode 100644 index dc2a0ed1..00000000 --- a/reme/memory/vector_based/base_memory_agent.py +++ /dev/null @@ -1,73 +0,0 @@ -"""Base memory agent for handling memory operations with tool-based reasoning.""" - -import json -from abc import ABCMeta - -from ...core.enumeration import MemoryType -from ...core.op import BaseReact -from ...core.schema import MemoryNode - - -class BaseMemoryAgent(BaseReact, metaclass=ABCMeta): - """Base class for memory agents that handle memory operations with tool-based reasoning.""" - - memory_type: MemoryType | None = None - - @property - def query(self) -> str: - """query""" - return self.context.get("query", "") - - @property - def messages(self) -> list: - """messages""" - return self.context.get("messages", []) - - @property - def description(self) -> str: - """description""" - return self.context.get("description", "") - - @property - def memory_target(self) -> str: - """memory_target""" - return self.context.memory_target - - @property - def history_node(self) -> MemoryNode: - """Returns the history node.""" - return self.context.history_node - - @property - def author(self) -> str: - """Returns the LLM model name as the author identifier.""" - return self.llm.model_name - - @property - def retrieved_nodes(self) -> list[MemoryNode]: - """Returns the retrieved nodes.""" - if "retrieved_nodes" not in self.context: - self.context.retrieved_nodes = [] - return self.context.retrieved_nodes - - @property - def memory_target_type_mapping(self) -> dict[str, MemoryType]: - """Get the memory target type mapping from context.""" - memory_targets = self.context.memory_targets - memory_target_type_mapping = self.context.service_context.memory_target_type_mapping.copy() - if memory_targets: - return {memory_target: memory_target_type_mapping[memory_target] for memory_target in memory_targets} - else: - return memory_target_type_mapping - - @property - def meta_memory_info(self) -> str: - """Get the meta memory info from context.""" - lines = [] - for memory_target, memory_type in self.memory_target_type_mapping.items(): - line = { - "agent": f"Agent managing {memory_type.value} memories for {memory_target}", - "memory_target": memory_target, - } - lines.append(json.dumps(line, ensure_ascii=False)) - return "\n".join(lines) diff --git a/reme/memory/vector_based/personal/__init__.py b/reme/memory/vector_based/personal/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/reme/memory/vector_based/personal/personal_retriever.py b/reme/memory/vector_based/personal/personal_retriever.py deleted file mode 100644 index 8b4965f1..00000000 --- a/reme/memory/vector_based/personal/personal_retriever.py +++ /dev/null @@ -1,148 +0,0 @@ -"""Personal memory retriever agent for retrieving personal memories through vector search.""" - -from loguru import logger - -from ..base_memory_agent import BaseMemoryAgent -from ....core.enumeration import MemoryType, Role -from ....core.op import BaseTool -from ....core.schema import Message -from ....core.utils import format_messages - -_PROFILE_TOOL_NAMES: tuple[str, ...] = ("retrieve_profile", "read_all_profiles") -_EMPTY_PROFILE_RESULTS: tuple[str, ...] = ("", "No profiles found.", "No new profiles found.") - - -class PersonalRetriever(BaseMemoryAgent): - """Retrieve personal memories through vector search and history reading.""" - - memory_type: MemoryType = MemoryType.PERSONAL - - def __init__(self, return_memory_nodes: bool = False, **kwargs): - super().__init__(**kwargs) - self.return_memory_nodes: bool = return_memory_nodes - - def _get_context(self) -> str: - if self.context.get("query"): - return self.context.query.strip() - if self.context.get("messages"): - return (self.description + "\n" + format_messages(self.context.messages)).strip() - raise ValueError("input must have either `query` or `messages`") - - async def _build_s1_messages(self, context: str) -> list[Message]: - return [ - Message( - role=Role.USER, - content=self.prompt_format( - prompt_name="user_message_s1", - memory_type=self.memory_type.value, - memory_target=self.memory_target, - context=context, - ), - ), - ] - - async def _build_s2_messages(self, context: str, profiles: str) -> list[Message]: - return [ - Message( - role=Role.USER, - content=self.prompt_format( - prompt_name="user_message_s2", - memory_type=self.memory_type.value, - memory_target=self.memory_target, - profiles=profiles, - context=context, - ), - ), - ] - - def _partition_tools(self) -> tuple[list[BaseTool], list[BaseTool]]: - profile_tools: list[BaseTool] = [] - memory_tools: list[BaseTool] = [] - for i, tool in enumerate(self.tools): - name = tool.tool_call.name - if name in _PROFILE_TOOL_NAMES: - profile_tools.append(tool) - else: - memory_tools.append(tool) - logger.info(f"[{self.__class__.__name__}] tool_call[{i}]={tool.tool_call.simple_input_dump(as_dict=False)}") - return profile_tools, memory_tools - - @staticmethod - def _extract_profile_context(tools: list[BaseTool]) -> str: - outputs = [] - for tool in tools: - response = getattr(tool, "response", None) - answer = getattr(response, "answer", "") - if answer and answer not in _EMPTY_PROFILE_RESULTS: - outputs.append(answer) - return "\n".join(outputs) - - async def _run_stage( - self, - stage: str, - messages: list[Message], - tools: list[BaseTool], - ) -> tuple[list[BaseTool], list[Message], bool]: - for message in messages: - role = message.name or message.role - logger.info(f"[{self.__class__.__name__} {stage}] role={role} {message.simple_dump(as_dict=False)}") - return await self.react(messages, tools, stage=stage) - - async def _acting_step( - self, - assistant_message: Message, - tools: list[BaseTool], - step: int, - stage: str = "", - **kwargs, - ) -> tuple[list[BaseTool], list[Message]]: - """Execute tool calls with memory context.""" - return await super()._acting_step( - assistant_message, - tools, - step, - memory_type=self.memory_type.value, - memory_target=self.memory_target, - retrieved_nodes=self.retrieved_nodes, - **kwargs, - ) - - async def execute(self): - context = self._get_context() - profile_tools, memory_tools = self._partition_tools() - - tools_s1: list[BaseTool] = [] - messages_s1: list[Message] = [] - success_s1 = True - profiles = "" - if profile_tools: - messages_s1 = await self._build_s1_messages(context) - tools_s1, messages_s1, success_s1 = await self._run_stage("s1-profile", messages_s1, profile_tools) - profiles = self._extract_profile_context(tools_s1) - - messages_s2 = await self._build_s2_messages(context, profiles) - tools_s2, messages_s2, success_s2 = await self._run_stage("s2-memory", messages_s2, memory_tools) - - answer = messages_s2[-1].content if success_s2 and messages_s2 else "" - result = { - "answer": answer, - "success": success_s1 and success_s2, - "messages": messages_s1 + messages_s2, - "tools": tools_s1 + tools_s2, - } - if self.return_memory_nodes: - result["answer"] = "\n".join( - [ - n.format( - include_memory_id=False, - include_when_to_use=False, - include_content=True, - include_message_time=True, - ref_memory_id_key="", - ) - for n in self.retrieved_nodes - ], - ) - - result["retrieved_nodes"] = self.retrieved_nodes - return result diff --git a/reme/memory/vector_based/personal/personal_summarizer.py b/reme/memory/vector_based/personal/personal_summarizer.py deleted file mode 100644 index 89b1316f..00000000 --- a/reme/memory/vector_based/personal/personal_summarizer.py +++ /dev/null @@ -1,136 +0,0 @@ -"""Personal memory summarizer agent for two-phase personal memory processing.""" - -from loguru import logger - -from ..base_memory_agent import BaseMemoryAgent -from ....core.enumeration import MemoryType, Role -from ....core.op import BaseTool -from ....core.schema import Message - -# Optional profile tools used to pre-load profile context; consumed by the -# summarizer itself and never exposed to the stage-two ReAct loop. -_PROFILE_CONTEXT_TOOLS: tuple[str, ...] = ("retrieve_profile", "read_all_profiles") - - -class PersonalSummarizer(BaseMemoryAgent): - """Two-phase personal memory processor: add memories, then update profiles.""" - - memory_type: MemoryType = MemoryType.PERSONAL - - async def _build_s1_messages(self) -> list[Message]: - return [ - Message( - role=Role.USER, - content=self.prompt_format( - prompt_name="user_message_s1", - context=self.context.history_node.content, - memory_type=self.memory_type.value, - memory_target=self.memory_target, - ), - ), - ] - - async def _build_s2_messages(self, profiles: str) -> list[Message]: - return [ - Message( - role=Role.USER, - content=self.prompt_format( - prompt_name="user_message_s2", - profiles=profiles, - context=self.context.history_node.content, - memory_type=self.memory_type.value, - memory_target=self.memory_target, - ), - ), - ] - - async def _acting_step( - self, - assistant_message: Message, - tools: list[BaseTool], - step: int, - stage: str = "", - **kwargs, - ) -> tuple[list[BaseTool], list[Message]]: - """Execute tool calls with memory context.""" - return await super()._acting_step( - assistant_message, - tools, - step, - stage=stage, - memory_type=self.memory_type.value, - memory_target=self.memory_target, - history_node=self.history_node, - author=self.author, - retrieved_nodes=self.retrieved_nodes, - **kwargs, - ) - - def _partition_tools(self) -> tuple[list[BaseTool], list[BaseTool], BaseTool | None]: - """Split attached tools into memory tools, profile tools, and a profile context tool.""" - memory_tools: list[BaseTool] = [] - profile_tools: list[BaseTool] = [] - profile_context_tool: BaseTool | None = None - for i, tool in enumerate(self.tools): - name = tool.tool_call.name - if name in _PROFILE_CONTEXT_TOOLS: - profile_context_tool = tool - elif "_memory" in name: - memory_tools.append(tool) - elif "_profile" in name: - profile_tools.append(tool) - else: - raise ValueError(f"[{self.__class__.__name__}] unknown tool_name={name}") - logger.info(f"[{self.__class__.__name__}] tool_call[{i}]={tool.tool_call.simple_input_dump(as_dict=False)}") - return memory_tools, profile_tools, profile_context_tool - - async def _preload_user_profile(self, tool: BaseTool | None) -> str: - """Invoke the profile context tool to obtain inline profile text.""" - if tool is None: - return "" - call_kwargs: dict = { - "memory_target": self.memory_target, - "service_context": self.service_context, - "retrieved_nodes": self.retrieved_nodes, - } - if tool.tool_call.name == "retrieve_profile": - call_kwargs["query"] = self.context.history_node.content - return await tool.call(**call_kwargs) - - async def _run_stage( - self, - stage: str, - messages: list[Message], - tools: list[BaseTool], - ) -> tuple[list[BaseTool], list[Message], bool]: - for message in messages: - role = message.name or message.role - logger.info(f"[{self.__class__.__name__} {stage}] role={role} {message.simple_dump(as_dict=False)}") - return await self.react(messages, tools, stage=stage) - - async def execute(self): - memory_tools, profile_tools, profile_context_tool = self._partition_tools() - - messages_s1 = await self._build_s1_messages() - tools_s1, messages_s1, success_s1 = await self._run_stage("s1-memory", messages_s1, memory_tools) - - if profile_tools: - profiles = await self._preload_user_profile(profile_context_tool) - messages_s2 = await self._build_s2_messages(profiles) - tools_s2, messages_s2, success_s2 = await self._run_stage("s2-profile", messages_s2, profile_tools) - else: - tools_s2, messages_s2, success_s2 = [], [], True - - answer = (messages_s1[-1].content if success_s1 and messages_s1 else "") + ( - messages_s2[-1].content if success_s2 and messages_s2 else "" - ) - tools = tools_s1 + tools_s2 - memory_nodes = [node for tool in tools for node in (tool.memory_nodes or [])] - - return { - "answer": answer, - "success": success_s1 and success_s2, - "messages": messages_s1 + messages_s2, - "tools": tools, - "memory_nodes": memory_nodes, - } diff --git a/reme/memory/vector_based/procedural/__init__.py b/reme/memory/vector_based/procedural/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/reme/memory/vector_based/procedural/procedural_retriever.py b/reme/memory/vector_based/procedural/procedural_retriever.py deleted file mode 100644 index a6265b62..00000000 --- a/reme/memory/vector_based/procedural/procedural_retriever.py +++ /dev/null @@ -1,82 +0,0 @@ -"""Procedural memory retriever agent for retrieving procedural memories through vector search.""" - -from ..base_memory_agent import BaseMemoryAgent -from ....core.enumeration import Role, MemoryType -from ....core.op import BaseTool -from ....core.schema import Message -from ....core.utils import format_messages - - -class ProceduralRetriever(BaseMemoryAgent): - """Retrieve procedural memories through vector search and history reading. - - Procedural memories represent "how-to" knowledge including: - - Workflows and step-by-step instructions - - Task execution patterns and best practices - - Success and failure patterns from past experiences - """ - - memory_type: MemoryType = MemoryType.PROCEDURAL - - def __init__(self, return_memory_nodes: bool = False, **kwargs): - super().__init__(**kwargs) - self.return_memory_nodes: bool = return_memory_nodes - - async def build_messages(self) -> list[Message]: - """Build messages with procedural memory retrieval context.""" - if self.context.get("query"): - context = self.context.query - elif self.context.get("messages"): - context = self.description + "\n" + format_messages(self.context.messages) - else: - raise ValueError("input must have either `query` or `messages`") - - return [ - Message( - role=Role.USER, - content=self.prompt_format( - prompt_name="user_message", - memory_type=self.memory_type.value, - memory_target=self.memory_target, - context=context.strip(), - ), - ), - ] - - async def _acting_step( - self, - assistant_message: Message, - tools: list[BaseTool], - step: int, - stage: str = "", - **kwargs, - ) -> tuple[list[BaseTool], list[Message]]: - """Execute tool calls with memory context.""" - return await super()._acting_step( - assistant_message, - tools, - step, - memory_type=self.memory_type.value, - memory_target=self.memory_target, - retrieved_nodes=self.retrieved_nodes, - **kwargs, - ) - - async def execute(self): - result = await super().execute() - if self.return_memory_nodes: - result["answer"] = "\n".join( - [ - n.format( - include_memory_id=False, - include_when_to_use=True, - include_content=True, - include_message_time=False, - ref_memory_id_key="", - ) - for n in self.retrieved_nodes - ], - ) - - result["retrieved_nodes"] = self.retrieved_nodes - return result diff --git a/reme/memory/vector_based/procedural/procedural_summarizer.py b/reme/memory/vector_based/procedural/procedural_summarizer.py deleted file mode 100644 index 25c485cb..00000000 --- a/reme/memory/vector_based/procedural/procedural_summarizer.py +++ /dev/null @@ -1,84 +0,0 @@ -"""Procedural memory summarizer agent for extracting and storing procedural knowledge.""" - -from loguru import logger - -from ..base_memory_agent import BaseMemoryAgent -from ....core.enumeration import Role, MemoryType -from ....core.op import BaseTool -from ....core.schema import Message - - -class ProceduralSummarizer(BaseMemoryAgent): - """Extract and store procedural memories from task execution trajectories. - - Procedural memories capture "how-to" knowledge including: - - Successful workflows and step-by-step approaches - - Lessons learned from failures and mistakes - - Best practices and optimization patterns - - Task execution strategies and techniques - """ - - memory_type: MemoryType = MemoryType.PROCEDURAL - - async def build_messages(self) -> list[Message]: - """Build messages for procedural memory extraction.""" - return [ - Message( - role=Role.USER, - content=self.prompt_format( - prompt_name="user_message", - context=self.context.history_node.content, - memory_type=self.memory_type.value, - memory_target=self.memory_target, - ), - ), - ] - - async def _acting_step( - self, - assistant_message: Message, - tools: list[BaseTool], - step: int, - stage: str = "", - **kwargs, - ) -> tuple[list[BaseTool], list[Message]]: - """Execute tool calls with memory context.""" - return await super()._acting_step( - assistant_message, - tools, - step, - stage=stage, - memory_type=self.memory_type.value, - memory_target=self.memory_target, - history_node=self.history_node, - author=self.author, - retrieved_nodes=self.retrieved_nodes, - **kwargs, - ) - - async def execute(self): - """Execute procedural memory extraction.""" - # Log available tools - for i, tool in enumerate(self.tools): - logger.info(f"[{self.__class__.__name__}] tool_call[{i}]={tool.tool_call.simple_input_dump(as_dict=False)}") - - messages = await self.build_messages() - for i, message in enumerate(messages): - role = message.name or message.role - logger.info(f"[{self.__class__.__name__}] role={role} {message.simple_dump(as_dict=False)}") - - tools, messages, success = await self.react(messages, self.tools) - - answer = messages[-1].content if success and messages else "" - memory_nodes = [] - for tool in tools: - if tool.memory_nodes: - memory_nodes.extend(tool.memory_nodes) - - return { - "answer": answer, - "success": success, - "messages": messages, - "tools": tools, - "memory_nodes": memory_nodes, - } diff --git a/reme/memory/vector_based/reme_retriever.py b/reme/memory/vector_based/reme_retriever.py deleted file mode 100644 index dcdade94..00000000 --- a/reme/memory/vector_based/reme_retriever.py +++ /dev/null @@ -1,94 +0,0 @@ -"""ReMe retriever agent that orchestrates multiple memory agents to retrieve information.""" - -from .base_memory_agent import BaseMemoryAgent -from ...core.enumeration import Role -from ...core.op import BaseTool -from ...core.schema import Message -from ...core.utils import format_messages - - -class ReMeRetriever(BaseMemoryAgent): - """Orchestrate multiple memory agents to retrieve information.""" - - async def build_messages(self) -> list[Message]: - if self.context.get("query"): - context = self.context.query - elif self.context.get("messages"): - context = self.description + "\n" + format_messages(self.context.messages) - else: - raise ValueError("input must have either `query` or `messages`") - - return [ - Message( - role=Role.SYSTEM, - content=self.prompt_format( - prompt_name="system_prompt", - meta_memory_info=self.meta_memory_info, - context=context.strip(), - ), - ), - Message( - role=Role.USER, - content=self.get_prompt("user_message"), - ), - ] - - async def _acting_step( - self, - assistant_message: Message, - tools: list[BaseTool], - step: int, - stage: str = "", - **kwargs, - ) -> tuple[list[BaseTool], list[Message]]: - return await super()._acting_step( - assistant_message, - tools, - step, - description=self.description, - messages=self.messages, - query=self.query, - author=self.author, - **kwargs, - ) - - async def react(self, messages: list[Message], tools: list["BaseTool"], stage: str = ""): - """Run single ReAct step - only one tool call iteration.""" - used_tools: list[BaseTool] = [] - assistant_message, should_act = await self._reasoning_step(messages, tools, step=0, stage=stage) - success = True - - if should_act: - t_tools, tool_messages = await self._acting_step(assistant_message, tools, step=0, stage=stage) - used_tools.extend(t_tools) - messages.extend(tool_messages) - - return used_tools, messages, success - - async def execute(self): - result = await super().execute() - tools: list[BaseTool] = result["tools"] - - answer = [] - success = True - messages = [] - tools_result = [] - retrieved_nodes = [] - - if tools: - delegate_task_tool = tools[0] - agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"] - for agent in agents: - answer.append(agent.response.answer) - success = success and agent.response.success - messages.extend(agent.response.metadata["messages"]) - tools_result.extend(agent.response.metadata["tools"]) - retrieved_nodes.extend(agent.response.metadata["retrieved_nodes"]) - - return { - "answer": "\n".join(answer), - "success": True, - "messages": messages, - "tools": tools_result, - "retrieved_nodes": retrieved_nodes, - } diff --git a/reme/memory/vector_based/reme_summarizer.py b/reme/memory/vector_based/reme_summarizer.py deleted file mode 100644 index 3687fe2c..00000000 --- a/reme/memory/vector_based/reme_summarizer.py +++ /dev/null @@ -1,98 +0,0 @@ -"""ReMe summarizer agent that orchestrates multiple memory agents to summarize information.""" - -from .base_memory_agent import BaseMemoryAgent -from ...core.enumeration import Role -from ...core.op import BaseTool -from ...core.schema import Message -from ...core.utils import format_messages - - -class ReMeSummarizer(BaseMemoryAgent): - """Orchestrates multiple memory agents to summarize and store information across different memory types.""" - - async def build_messages(self) -> list[Message]: - add_history_tool: BaseTool | None = self.pop_tool("add_history") - if add_history_tool is not None: - await add_history_tool.call( - messages=self.messages, - description=self.description, - author=self.author, - service_context=self.service_context, - ) - self.context.history_node = add_history_tool.context.history_node - - context = self.context.description + "\n" + format_messages(self.context.messages) - messages = [ - Message( - role=Role.SYSTEM, - content=self.prompt_format( - prompt_name="system_prompt", - meta_memory_info=self.meta_memory_info, - context=context.strip(), - ), - ), - Message( - role=Role.USER, - content=self.get_prompt("user_message"), - ), - ] - - return messages - - async def _acting_step( - self, - assistant_message: Message, - tools: list[BaseTool], - step: int, - stage: str = "", - **kwargs, - ) -> tuple[list[BaseTool], list[Message]]: - return await super()._acting_step( - assistant_message, - tools, - step, - description=self.description, - messages=self.messages, - history_node=self.history_node, - author=self.author, - **kwargs, - ) - - async def react(self, messages: list[Message], tools: list["BaseTool"], stage: str = ""): - """Run single ReAct step - only one tool call iteration.""" - used_tools: list[BaseTool] = [] - assistant_message, should_act = await self._reasoning_step(messages, tools, step=0, stage=stage) - success = True - - if should_act: - t_tools, tool_messages = await self._acting_step(assistant_message, tools, step=0, stage=stage) - used_tools.extend(t_tools) - messages.extend(tool_messages) - - return used_tools, messages, success - - async def execute(self): - result = await super().execute() - tools: list[BaseTool] = result["tools"] - - success = True - messages = [] - tools_result = [] - memory_nodes = [] - - if tools: - delegate_task_tool = tools[0] - agents: list[BaseMemoryAgent] = delegate_task_tool.response.metadata["agents"] - - for agent in agents: - success = success and agent.response.success - messages.extend(agent.response.metadata["messages"]) - tools_result.extend(agent.response.metadata["tools"]) - memory_nodes.extend(agent.response.metadata["memory_nodes"]) - - return { - "answer": memory_nodes, - "success": True, - "messages": messages, - "tools": tools_result, - } diff --git a/reme/memory/vector_based/tool_call/__init__.py b/reme/memory/vector_based/tool_call/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/reme/memory/vector_based/tool_call/tool_retriever.py b/reme/memory/vector_based/tool_call/tool_retriever.py deleted file mode 100644 index 722f9720..00000000 --- a/reme/memory/vector_based/tool_call/tool_retriever.py +++ /dev/null @@ -1,83 +0,0 @@ -"""Tool memory retriever agent for retrieving tool usage experiences through vector search.""" - -from ..base_memory_agent import BaseMemoryAgent -from ....core.enumeration import Role, MemoryType -from ....core.op import BaseTool -from ....core.schema import Message -from ....core.utils import format_messages - - -class ToolRetriever(BaseMemoryAgent): - """Retrieve tool memories through vector search and history reading. - - Tool memories represent knowledge about tool usage including: - - Successful tool invocations and their parameters - - Failed tool calls and lessons learned - - Tool selection strategies for different scenarios - - Parameter optimization patterns - """ - - memory_type: MemoryType = MemoryType.TOOL - - def __init__(self, return_memory_nodes: bool = False, **kwargs): - super().__init__(**kwargs) - self.return_memory_nodes: bool = return_memory_nodes - - async def build_messages(self) -> list[Message]: - """Build messages with tool memory retrieval context.""" - if self.context.get("query"): - context = self.context.query - elif self.context.get("messages"): - context = self.description + "\n" + format_messages(self.context.messages) - else: - raise ValueError("input must have either `query` or `messages`") - - return [ - Message( - role=Role.USER, - content=self.prompt_format( - prompt_name="user_message", - memory_type=self.memory_type.value, - memory_target=self.memory_target, - context=context.strip(), - ), - ), - ] - - async def _acting_step( - self, - assistant_message: Message, - tools: list[BaseTool], - step: int, - stage: str = "", - **kwargs, - ) -> tuple[list[BaseTool], list[Message]]: - """Execute tool calls with memory context.""" - return await super()._acting_step( - assistant_message, - tools, - step, - memory_type=self.memory_type.value, - memory_target=self.memory_target, - retrieved_nodes=self.retrieved_nodes, - **kwargs, - ) - - async def execute(self): - result = await super().execute() - if self.return_memory_nodes: - result["answer"] = "\n".join( - [ - n.format( - include_memory_id=False, - include_when_to_use=True, - include_content=True, - include_message_time=False, - ref_memory_id_key="", - ) - for n in self.retrieved_nodes - ], - ) - - result["retrieved_nodes"] = self.retrieved_nodes - return result diff --git a/reme/memory/vector_based/tool_call/tool_summarizer.py b/reme/memory/vector_based/tool_call/tool_summarizer.py deleted file mode 100644 index a8a823ca..00000000 --- a/reme/memory/vector_based/tool_call/tool_summarizer.py +++ /dev/null @@ -1,84 +0,0 @@ -"""Tool memory summarizer agent for extracting and storing tool usage experiences.""" - -from loguru import logger - -from ..base_memory_agent import BaseMemoryAgent -from ....core.enumeration import Role, MemoryType -from ....core.op import BaseTool -from ....core.schema import Message - - -class ToolSummarizer(BaseMemoryAgent): - """Extract and store tool memories from task execution trajectories. - - Tool memories capture knowledge about tool usage including: - - Successful tool invocations with effective parameters - - Failed tool calls and why they failed - - Tool selection strategies for different scenarios - - Parameter optimization insights - """ - - memory_type: MemoryType = MemoryType.TOOL - - async def build_messages(self) -> list[Message]: - """Build messages for tool memory extraction.""" - return [ - Message( - role=Role.USER, - content=self.prompt_format( - prompt_name="user_message", - context=self.context.history_node.content, - memory_type=self.memory_type.value, - memory_target=self.memory_target, - ), - ), - ] - - async def _acting_step( - self, - assistant_message: Message, - tools: list[BaseTool], - step: int, - stage: str = "", - **kwargs, - ) -> tuple[list[BaseTool], list[Message]]: - """Execute tool calls with memory context.""" - return await super()._acting_step( - assistant_message, - tools, - step, - stage=stage, - memory_type=self.memory_type.value, - memory_target=self.memory_target, - history_node=self.history_node, - author=self.author, - retrieved_nodes=self.retrieved_nodes, - **kwargs, - ) - - async def execute(self): - """Execute tool memory extraction.""" - # Log available tools - for i, tool in enumerate(self.tools): - logger.info(f"[{self.__class__.__name__}] tool_call[{i}]={tool.tool_call.simple_input_dump(as_dict=False)}") - - messages = await self.build_messages() - for i, message in enumerate(messages): - role = message.name or message.role - logger.info(f"[{self.__class__.__name__}] role={role} {message.simple_dump(as_dict=False)}") - - tools, messages, success = await self.react(messages, self.tools) - - answer = messages[-1].content if success and messages else "" - memory_nodes = [] - for tool in tools: - if tool.memory_nodes: - memory_nodes.extend(tool.memory_nodes) - - return { - "answer": answer, - "success": success, - "messages": messages, - "tools": tools, - "memory_nodes": memory_nodes, - } diff --git a/reme/memory/vector_tools/__init__.py b/reme/memory/vector_tools/__init__.py deleted file mode 100644 index 9fb04923..00000000 --- a/reme/memory/vector_tools/__init__.py +++ /dev/null @@ -1,66 +0,0 @@ -"""memory tools""" - -# pylint: disable=no-name-in-module - -from .base_memory_tool import BaseMemoryTool - -# chunk tools -from .delegate_task import DelegateTask - -# history tools -from .history.add_history import AddHistory -from .history.read_history import ReadHistory -from .history.read_history_v2 import ReadHistoryV2 - -# profiles tools -from .profiles.add_draft_and_read_all_profiles import AddDraftAndReadAllProfiles -from .profiles.add_profile import AddProfile -from .profiles.delete_profile import DeleteProfile -from .profiles.read_all_profiles import ReadAllProfiles -from .profiles.retrieve_profile import RetrieveProfile -from .profiles.update_profile import UpdateProfile -from .profiles.update_profiles_v1 import UpdateProfilesV1 - -# record tools -from .record.add_and_retrieve_similar_memory import AddAndRetrieveSimilarMemory -from .record.add_draft_and_retrieve_similar_memory import AddDraftAndRetrieveSimilarMemory -from .record.add_memory import AddMemory -from .record.delete_memory import DeleteMemory -from .record.retrieve_memory import RetrieveMemory -from .record.retrieve_recent_memory import RetrieveRecentMemory -from .record.update_memory import UpdateMemory -from .record.update_memory_v1 import UpdateMemoryV1 -from .record.update_memory_v2 import UpdateMemoryV2 -from ...core import R - -__all__ = [ - # base - "BaseMemoryTool", - "DelegateTask", - # history tools - "AddHistory", - "ReadHistory", - "ReadHistoryV2", - # profiles tools - "AddDraftAndReadAllProfiles", - "AddProfile", - "DeleteProfile", - "ReadAllProfiles", - "RetrieveProfile", - "UpdateProfile", - "UpdateProfilesV1", - # record tools - "AddAndRetrieveSimilarMemory", - "AddDraftAndRetrieveSimilarMemory", - "AddMemory", - "DeleteMemory", - "RetrieveMemory", - "RetrieveRecentMemory", - "UpdateMemory", - "UpdateMemoryV1", - "UpdateMemoryV2", -] - -for name in __all__: - tool_class = globals()[name] - R.ops.register(tool_class) diff --git a/reme/memory/vector_tools/base_memory_tool.py b/reme/memory/vector_tools/base_memory_tool.py deleted file mode 100644 index f0781bd3..00000000 --- a/reme/memory/vector_tools/base_memory_tool.py +++ /dev/null @@ -1,128 +0,0 @@ -"""Base class for memory tool""" - -from abc import ABCMeta -from pathlib import Path - -from .profiles.profile_handler import ProfileHandler -from ...core.enumeration import MemoryType -from ...core.op import BaseTool -from ...core.schema import ToolCall, MemoryNode, ToolAttr - - -class BaseMemoryTool(BaseTool, metaclass=ABCMeta): - """Base class for memory tool""" - - def __init__( - self, - enable_multiple: bool = True, - enable_thinking_params: bool = False, - profile_dir: str = "", - profile_backend: str = "filesystem", - profile_store_name: str = "profile", - profile_max_capacity: int = 50, - **kwargs, - ): - super().__init__(**kwargs) - self.enable_multiple: bool = enable_multiple - self.enable_thinking_params: bool = enable_thinking_params - self.profile_dir: str = profile_dir - self.profile_backend: str = profile_backend - self.profile_store_name: str = profile_store_name - self.profile_max_capacity: int = profile_max_capacity - - def _build_tool_call(self) -> ToolCall: - """Build and return the tool call schema""" - - def _build_multiple_tool_call(self) -> ToolCall: - """Build and return the multiple tool call schema""" - - @property - def tool_call(self) -> ToolCall | None: - """Get the tool call schema.""" - if self._tool_call is None: - if self.enable_multiple: - self._tool_call = self._build_multiple_tool_call() - else: - self._tool_call = self._build_tool_call() - self._tool_call.name = self._tool_call.name or self.name - - # Add thinking parameter if enabled - if self.enable_thinking_params: - parameters = self._tool_call.parameters - if parameters and parameters.properties is not None: - if "thinking" not in parameters.properties: - parameters.properties = { - "thinking": ToolAttr( - type="string", - description="Your complete and detailed thinking process " - "about how to fill in each parameter", - ), - **parameters.properties, - } - if parameters.required is not None: - parameters.required = ["thinking", *parameters.required] - else: - parameters.required = ["thinking"] - return self._tool_call - - @property - def memory_type(self) -> MemoryType: - """Get the memory type from context.""" - return self.memory_target_type_mapping[self.memory_target] - - @property - def memory_target(self) -> str: - """Get the memory target from context.""" - if "memory_target" in self.context: - return self.context.memory_target - elif len(self.memory_target_type_mapping) == 1: - return list(self.memory_target_type_mapping.keys())[0] - else: - raise ValueError("memory_target is not specified in context or memory_target_type_mapping!") - - @property - def history_id(self) -> str: - """Get the history node from context.""" - if "history_node" in self.context: - return self.context.history_node.memory_id - return "" - - @property - def retrieved_nodes(self) -> list[MemoryNode]: - """Get the retrieved nodes from context.""" - return self.context.retrieved_nodes - - @property - def author(self) -> str: - """Get the author from context.""" - return self.context.author - - @property - def memory_nodes(self) -> list[MemoryNode | str]: - """Get the memory nodes from context.""" - if "memory_nodes" not in self.context: - self.context.memory_nodes = [] - return self.context.memory_nodes - - @property - def memory_target_type_mapping(self) -> dict[str, MemoryType]: - """Get the memory target type mapping from context.""" - return self.context.service_context.memory_target_type_mapping - - @property - def profile_path(self) -> Path | None: - """Get the path to the profile directory for the current collection.""" - if not self.profile_dir: - return None - return Path(self.profile_dir) / self.vector_store.collection_name - - def get_profile_handler(self, memory_target: str) -> ProfileHandler: - """Build a profile handler for the current backend configuration.""" - return ProfileHandler( - memory_target=memory_target, - profile_path=self.profile_path, - service_context=self.service_context, - profile_backend=self.profile_backend, - profile_store_name=self.profile_store_name, - max_capacity=self.profile_max_capacity, - ) diff --git a/reme/memory/vector_tools/delegate_task.py b/reme/memory/vector_tools/delegate_task.py deleted file mode 100644 index d3a0298d..00000000 --- a/reme/memory/vector_tools/delegate_task.py +++ /dev/null @@ -1,89 +0,0 @@ -"""Hands-off tool to delegate memory tasks to specific agents""" - -from loguru import logger - -from .base_memory_tool import BaseMemoryTool -from ..vector_based import BaseMemoryAgent -from ...core.enumeration import MemoryType -from ...core.schema import ToolCall - - -class DelegateTask(BaseMemoryTool): - """Tool to delegate memory tasks to appropriate memory agents""" - - def __init__(self, memory_agents: list[BaseMemoryAgent] = None, **kwargs): - kwargs["enable_multiple"] = True - kwargs["sub_ops"] = memory_agents or [] - super().__init__(**kwargs) - self.sub_ops: list[BaseMemoryAgent] = [a for a in self.sub_ops if isinstance(a, BaseMemoryAgent)] - assert all(a.memory_type is not None for a in self.sub_ops) - - @property - def memory_agent_dict(self) -> dict[MemoryType, BaseMemoryAgent]: - """Map memory types to their corresponding agents""" - return {a.memory_type: a for a in self.sub_ops} - - def _build_multiple_tool_call(self) -> ToolCall: - """Build and return the multiple tool call schema""" - return ToolCall( - **{ - "description": "Delegate tasks to appropriate agents.", - "parameters": { - "type": "object", - "properties": { - "tasks": { - "type": "array", - "description": "List of tasks to delegate to specific memory agents", - "items": { - "type": "object", - "description": "A task item", - "properties": { - "memory_target": { - "type": "string", - "description": "The memory_target identifier to " - "delegate to the corresponding agent", - }, - }, - "required": ["memory_target"], - }, - }, - }, - "required": ["tasks"], - }, - }, - ) - - async def execute(self): - # Deduplicate and validate tasks - tasks = self.context.get("tasks", []) - memory_target_tasks = sorted(set(task["memory_target"] for task in tasks)) - - # Submit memory_target_tasks to agents - agent_list: list[BaseMemoryAgent] = [] - for i, memory_target in enumerate(memory_target_tasks): - if memory_target not in self.memory_target_type_mapping: - logger.warning(f"Memory target {memory_target} not found in memory_target_type_mapping") - continue - - memory_type = self.memory_target_type_mapping[memory_target] - agent = self.memory_agent_dict[memory_type].copy() - agent_list.append(agent) - - logger.info(f"Task {i}: {memory_type.value} agent for {memory_target}") - task_kwargs = {"memory_target": memory_target} - for k in ["query", "messages", "description", "history_node"]: - if k in self.context: - task_kwargs[k] = self.context[k] - self.submit_async_task(agent.call, service_context=self.service_context, **task_kwargs) - await self.join_async_tasks() - - # Collect results - results = [] - for agent in agent_list: - results.append(f"Task: {agent.memory_target}\n{agent.response.answer}") - - logger.info(f"Completed {len(results)} memory_target(s)") - return { - "answer": "\n\n".join(results), - "agents": agent_list, - } diff --git a/reme/memory/vector_tools/history/__init__.py b/reme/memory/vector_tools/history/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/reme/memory/vector_tools/history/add_history.py b/reme/memory/vector_tools/history/add_history.py deleted file mode 100644 index 24fb3228..00000000 --- a/reme/memory/vector_tools/history/add_history.py +++ /dev/null @@ -1,57 +0,0 @@ -"""Add history tool""" - -import json - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ....core.enumeration import MemoryType -from ....core.schema import ToolCall, MemoryNode, Message -from ....core.utils import format_messages - - -class AddHistory(BaseMemoryTool): - """Tool to add historical dialogue to vector store""" - - def __init__(self, **kwargs): - kwargs["enable_multiple"] = False - super().__init__(**kwargs) - - def _build_tool_call(self) -> ToolCall: - """Build and return the tool call schema""" - return ToolCall( - **{ - "description": "Add original history dialogue.", - "parameters": { - "type": "object", - "properties": {}, - "required": [], - }, - }, - ) - - async def execute(self): - """Execute the add history operation""" - self.context.messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] - history_content: str = self.context.description + "\n" + format_messages(self.context.messages) - history_content = history_content.strip() - history_node = MemoryNode( - memory_type=MemoryType.HISTORY, - when_to_use=history_content[:1024], - content=history_content, - author=self.author, - metadata={ - "messages": json.dumps( - [m.model_dump(exclude_none=True) for m in self.context.messages], - ensure_ascii=False, - ), - }, - ) - self.context.history_node = history_node - logger.info(f"Adding history node: {history_node.model_dump_json(indent=2)}") - - vector_node = history_node.to_vector_node() - await self.vector_store.delete(vector_node.vector_id) - await self.vector_store.insert([vector_node]) - - return f"Successfully added history: {history_node.memory_id}" diff --git a/reme/memory/vector_tools/history/read_history.py b/reme/memory/vector_tools/history/read_history.py deleted file mode 100644 index efa7ac3b..00000000 --- a/reme/memory/vector_tools/history/read_history.py +++ /dev/null @@ -1,77 +0,0 @@ -"""Read history memory tool""" - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import MemoryNode, ToolCall - - -class ReadHistory(BaseMemoryTool): - """Read history memory tool""" - - def __init__(self, **kwargs): - kwargs.setdefault("enable_multiple", False) - super().__init__(**kwargs) - - def _build_tool_call(self) -> ToolCall: - """Build and return the tool call schema""" - return ToolCall( - **{ - "description": "Read original history dialogue.", - "parameters": { - "type": "object", - "properties": { - "history_id": { - "type": "string", - "description": "history_id", - }, - }, - "required": ["history_id"], - }, - }, - ) - - def _build_multiple_tool_call(self) -> ToolCall: - """Build and return the tool call schema for multiple histories""" - return ToolCall( - **{ - "description": "Read multiple original history dialogues by their IDs.", - "parameters": { - "type": "object", - "properties": { - "history_ids": { - "type": "array", - "items": {"type": "string"}, - "description": "List of history IDs to read", - }, - }, - "required": ["history_ids"], - }, - }, - ) - - async def execute(self): - """Execute the tool call""" - history_ids = self.context.history_ids if self.enable_multiple else [self.context.history_id] - - if not history_ids or (len(history_ids) == 1 and not history_ids[0]): - output = "No history_ids provided." - logger.warning(output) - return output - - nodes = await self.vector_store.get(vector_ids=history_ids) - - if not nodes: - output = f"No data found for history_ids={history_ids}." - logger.warning(output) - return output - - results = [] - for node in nodes: - memory_node: MemoryNode = MemoryNode.from_vector_node(node) - self.retrieved_nodes.append(memory_node) - results.append(f"Historical Dialogue[{memory_node.memory_id}]\n{memory_node.content}") - - output = "\n\n".join(results) if self.enable_multiple else results[0] - logger.info(f"Successfully read {len(nodes)} history memory_node(s): {history_ids}") - return output diff --git a/reme/memory/vector_tools/history/read_history_v2.py b/reme/memory/vector_tools/history/read_history_v2.py deleted file mode 100644 index 7ca267e6..00000000 --- a/reme/memory/vector_tools/history/read_history_v2.py +++ /dev/null @@ -1,123 +0,0 @@ -"""Read history memory tool""" - -import json - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import MemoryNode, ToolCall, Message -from ....core.utils import format_messages, cosine_similarity - - -class ReadHistoryV2(BaseMemoryTool): - """Read history memory tool""" - - def __init__(self, message_block_size: int = 4, vector_top_k: int = 5, **kwargs): - kwargs.setdefault("enable_multiple", False) - super().__init__(name="read_history", **kwargs) - self.message_block_size: int = message_block_size - self.vector_top_k: int = vector_top_k - - def _build_tool_call(self) -> ToolCall: - """Build and return the tool call schema""" - return ToolCall( - **{ - "description": "Read original history dialogue.", - "parameters": { - "type": "object", - "properties": { - "history_id": { - "type": "string", - "description": "history_id", - }, - "query": { - "type": "string", - "description": "Query to filter the history", - }, - }, - "required": ["history_id", "query"], - }, - }, - ) - - def _build_multiple_tool_call(self) -> ToolCall: - """Build and return the tool call schema for multiple histories""" - return ToolCall( - **{ - "description": "Read multiple original history dialogues by their IDs.", - "parameters": { - "type": "object", - "properties": { - "history_items": { - "type": "array", - "items": { - "type": "object", - "properties": { - "history_id": { - "type": "string", - "description": "history_id", - }, - "query": { - "type": "string", - "description": "Query to filter the history", - }, - }, - "required": ["history_id", "query"], - }, - "description": "List of history items to read, each containing history_id and query", - }, - }, - "required": ["history_items"], - }, - }, - ) - - async def execute(self): - """Execute the tool call""" - if self.enable_multiple: - history_items = self.context.history_items - else: - history_items = [{"history_id": self.context.history_id, "query": self.context.query}] - - if not history_items: - output = "No history_ids provided." - logger.warning(output) - return output - - all_results = [] - for history_item in history_items: - history_id = history_item["history_id"] - query = history_item["query"] - query_embedding = await self.embedding_model.get_embedding(query) - - history_node: MemoryNode = await self.vector_store.get(vector_ids=history_id) - messages = json.loads(history_node.metadata["messages"]) - messages = [Message(**m) for m in messages] - - message_blocks = [] - for i in range(0, len(messages), self.message_block_size): - block = messages[i : i + self.message_block_size] - message_blocks.append(block) - - block_similarities = [] - for block in message_blocks: - block_text = format_messages(block, add_index=False) - block_embedding = await self.embedding_model.get_embedding(block_text) - similarity = cosine_similarity(query_embedding, block_embedding) - block_similarities.append((similarity, block_text)) - - block_similarities.sort(key=lambda x: x[0], reverse=True) - top_k_blocks = block_similarities[: self.vector_top_k] - result_text = "\n".join([block_text for _, block_text in top_k_blocks]) - all_results.append(result_text) - - if not all_results: - history_ids = [item["history_id"] for item in history_items] - output = f"No data found for history_ids={history_ids}." - logger.warning(output) - return output - - output = "\n\n".join(all_results) - history_ids = [item["history_id"] for item in history_items] - logger.info(f"Successfully read {len(all_results)} history result(s): {history_ids}") - return output diff --git a/reme/memory/vector_tools/profiles/__init__.py b/reme/memory/vector_tools/profiles/__init__.py deleted file mode 100644 index 150455db..00000000 --- a/reme/memory/vector_tools/profiles/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Profile memory tools.""" diff --git a/reme/memory/vector_tools/profiles/add_draft_and_read_all_profiles.py b/reme/memory/vector_tools/profiles/add_draft_and_read_all_profiles.py deleted file mode 100644 index c8e4889a..00000000 --- a/reme/memory/vector_tools/profiles/add_draft_and_read_all_profiles.py +++ /dev/null @@ -1,106 +0,0 @@ -"""Add draft profile and read all profiles from the configured backend.""" - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall - - -class AddDraftAndReadAllProfiles(BaseMemoryTool): - """Tool to add draft profile and read all profiles""" - - def __init__(self, enable_memory_target: bool = False, **kwargs): - super().__init__(**kwargs) - self.enable_memory_target: bool = enable_memory_target - - def _build_query_parameters(self) -> dict: - """Build the query parameters schema""" - properties = { - "message_time": { - "type": "string", - "description": "Message time, e.g. '2020-01-01 00:00:00'", - }, - "profile_key": { - "type": "string", - "description": "Profile key or category, e.g. 'name'", - }, - "profile_value": { - "type": "string", - "description": "Profile value or content, e.g. 'John Smith'", - }, - } - required = ["message_time", "profile_key", "profile_value"] - - if self.enable_memory_target: - properties["memory_target"] = { - "type": "string", - "description": "memory_target", - } - required.append("memory_target") - - return { - "type": "object", - "properties": properties, - "required": required, - } - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "Add draft profile and read all profiles from local storage.", - "parameters": self._build_query_parameters(), - }, - ) - - def _build_multiple_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "Add draft profile and read all profiles from local storage.", - "parameters": { - "type": "object", - "properties": { - "draft_items": { - "type": "array", - "description": "List of draft profile items.", - "items": self._build_query_parameters(), - }, - }, - "required": ["draft_items"], - }, - }, - ) - - async def execute(self): - if self.enable_multiple: - draft_items = self.context.get("draft_items", []) - else: - draft_items = [self.context] - - # Collect all profiles from all targets - all_profiles = [] - targets_processed = set() - - for item in draft_items: - if self.enable_memory_target: - target = item["memory_target"] - else: - target = self.memory_target - - # Skip if already processed this target - if target in targets_processed: - continue - targets_processed.add(target) - - profile_handler = self.get_profile_handler(target) - profiles_str = await profile_handler.aread_all(add_profile_id=True) - if profiles_str: - all_profiles.append(f"## Profiles for {target}:\n{profiles_str}") - - if not all_profiles: - output = "No profiles found." - logger.info(output) - return output - - output = "\n\n".join(all_profiles) - logger.info(f"Successfully read profiles for {len(targets_processed)} target(s)") - return output diff --git a/reme/memory/vector_tools/profiles/add_profile.py b/reme/memory/vector_tools/profiles/add_profile.py deleted file mode 100644 index 5654ebb7..00000000 --- a/reme/memory/vector_tools/profiles/add_profile.py +++ /dev/null @@ -1,70 +0,0 @@ -"""Add user profile tool.""" - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall - - -class AddProfile(BaseMemoryTool): - """Tool to add a single profile entry""" - - def __init__(self, **kwargs): - super().__init__(**kwargs) - - def _build_tool_call(self) -> ToolCall: - """Build and return the tool call schema""" - return ToolCall( - **{ - "description": "Add a new profile entry for the user.", - "parameters": { - "type": "object", - "properties": { - "message_time": { - "type": "string", - "description": "Message time, e.g. '2020-01-01 00:00:00'", - }, - "profile_key": { - "type": "string", - "description": "Profile key or category, e.g. 'name'", - }, - "profile_value": { - "type": "string", - "description": "Profile value or content, e.g. 'John Smith'", - }, - }, - "required": ["message_time", "profile_key", "profile_value"], - }, - }, - ) - - async def execute(self): - profile_handler = self.get_profile_handler(self.memory_target) - - # Get parameters - message_time = self.context.get("message_time", "") - profile_key = self.context.get("profile_key", "") - profile_value = self.context.get("profile_value", "") - - if not profile_key or not profile_value: - return "Missing required parameters (profile_key or profile_value), operation cancelled." - - # Build profile dict - profile = { - "message_time": message_time, - "profile_key": profile_key, - "profile_value": profile_value, - } - - # Add profile using ProfileHandler - new_nodes = await profile_handler.aadd_batch(profiles=[profile], ref_memory_id=self.history_id) - self.memory_nodes.extend(new_nodes) - - if new_nodes: - output = f"Successfully added profile: [{profile_key}] = {profile_value}" - logger.info(output) - return output - else: - output = "Failed to add profile." - logger.warning(output) - return output diff --git a/reme/memory/vector_tools/profiles/delete_profile.py b/reme/memory/vector_tools/profiles/delete_profile.py deleted file mode 100644 index 1c5f845b..00000000 --- a/reme/memory/vector_tools/profiles/delete_profile.py +++ /dev/null @@ -1,52 +0,0 @@ -"""Delete user profile tool.""" - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall - - -class DeleteProfile(BaseMemoryTool): - """Tool to delete a single profile entry by ID""" - - def __init__(self, **kwargs): - super().__init__(**kwargs) - - def _build_tool_call(self) -> ToolCall: - """Build and return the tool call schema""" - return ToolCall( - **{ - "description": "Delete a profile entry by profile ID.", - "parameters": { - "type": "object", - "properties": { - "profile_id": { - "type": "string", - "description": "The unique ID of the profile to delete.", - }, - }, - "required": ["profile_id"], - }, - }, - ) - - async def execute(self): - profile_handler = self.get_profile_handler(self.memory_target) - - # Get profile_id parameter - profile_id = self.context.get("profile_id", "") - - if not profile_id: - return "No profile_id provided, operation cancelled." - - # Delete profile using ProfileHandler - success = await profile_handler.adelete(profile_id) - - if success: - output = f"Successfully deleted profile with ID: {profile_id}" - logger.info(output) - return output - else: - output = f"Profile with ID '{profile_id}' not found." - logger.warning(output) - return output diff --git a/reme/memory/vector_tools/profiles/file_profile_backend.py b/reme/memory/vector_tools/profiles/file_profile_backend.py deleted file mode 100644 index e6226459..00000000 --- a/reme/memory/vector_tools/profiles/file_profile_backend.py +++ /dev/null @@ -1,234 +0,0 @@ -"""Filesystem-backed profile storage.""" - -from pathlib import Path - -from loguru import logger - -from .profile_backend import BaseProfileBackend -from ....core.enumeration import MemoryType -from ....core.schema import MemoryNode -from ....core.utils import CacheHandler, deduplicate_memories - - -class FileProfileBackend(BaseProfileBackend): - """Persist user profiles in local JSONL cache files.""" - - def __init__(self, profile_path: str | Path, memory_target: str, max_capacity: int = 50): - super().__init__(memory_target=memory_target, max_capacity=max_capacity) - self.cache_key: str = self.memory_target.replace(" ", "_").lower() - self.cache_handler: CacheHandler = CacheHandler(profile_path) - - def _load_nodes(self) -> list[MemoryNode]: - cached_data = self.cache_handler.load(self.cache_key, auto_clean=False) - if not cached_data: - return [] - return [MemoryNode(**data) for data in cached_data] - - def _save_nodes(self, nodes: list[MemoryNode], apply_limits: bool = True): - if apply_limits: - nodes = deduplicate_memories(nodes) - - if len(nodes) > self.max_capacity: - sorted_nodes = sorted(nodes, key=lambda n: n.message_time) - removed_count = len(sorted_nodes) - self.max_capacity - nodes = sorted_nodes[removed_count:] - logger.info( - f"Capacity limit reached: removed {removed_count} oldest profiles " - f"(kept {len(nodes)}/{self.max_capacity})", - ) - - nodes_data = [node.model_dump(exclude_none=True) for node in nodes] - self.cache_handler.save(self.cache_key, nodes_data) - logger.info(f"Saved {len(nodes)} profiles to {self.cache_key}") - - def get_all_sync(self) -> list[MemoryNode]: - """Load all profile nodes from cache, ordered by ``message_time``.""" - nodes = self._load_nodes() - nodes.sort(key=lambda n: n.message_time) - return nodes - - def get_by_sync(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None: - """Return the first node matching ``profile_id`` or ``profile_key``.""" - if not profile_id and not profile_key: - raise ValueError("Must provide either profile_id or profile_key") - - for node in self._load_nodes(): - if profile_id and node.memory_id == profile_id: - return node - if profile_key and node.when_to_use == profile_key: - return node - return None - - def delete_sync(self, profile_id: str | list[str]) -> bool | int: - """Remove one id, many ids, or none; returns bool, count, or 0/false if nothing removed.""" - nodes = self._load_nodes() - original_count = len(nodes) - - if isinstance(profile_id, list): - profile_ids_set = set(profile_id) - nodes = [n for n in nodes if n.memory_id not in profile_ids_set] - deleted_count = original_count - len(nodes) - if deleted_count == 0: - logger.warning(f"No profiles found to delete from {len(profile_id)} IDs") - return 0 - - self._save_nodes(nodes, apply_limits=False) - logger.info(f"Batch deleted {deleted_count} profiles") - return deleted_count - - nodes = [n for n in nodes if n.memory_id != profile_id] - if len(nodes) == original_count: - logger.warning(f"Profile {profile_id} not found") - return False - - self._save_nodes(nodes, apply_limits=False) - logger.info(f"Deleted profile {profile_id}") - return True - - def delete_all_sync(self) -> int: - """Clear every cached profile for this target; returns how many were stored.""" - nodes = self._load_nodes() - count = len(nodes) - self._save_nodes([], apply_limits=False) - logger.info(f"Deleted all {count} profiles") - return count - - def add_sync(self, message_time: str, profile_key: str, profile_value: str, ref_memory_id: str = "") -> MemoryNode: - """Append a profile row, replacing any existing row with the same key.""" - nodes = self._load_nodes() - - new_node = MemoryNode( - memory_type=MemoryType.PERSONAL, - memory_target=self.memory_target, - when_to_use=profile_key, - content=profile_value, - message_time=message_time, - ref_memory_id=ref_memory_id, - ) - - original_count = len(nodes) - nodes = [n for n in nodes if n.when_to_use != profile_key] - if len(nodes) < original_count: - logger.info(f"Removed {original_count - len(nodes)} duplicate profile(s) with key: {profile_key}") - - nodes.append(new_node) - self._save_nodes(nodes) - logger.info(f"Added profile: {profile_key}={profile_value}") - return new_node - - def add_batch_sync(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]: - """Insert many profiles in one write, deduping by key against existing rows.""" - if not profiles: - return [] - - nodes = self._load_nodes() - new_nodes = [ - MemoryNode( - memory_type=MemoryType.PERSONAL, - memory_target=self.memory_target, - when_to_use=p.get("profile_key", ""), - content=p.get("profile_value", ""), - message_time=p.get("message_time", ""), - ref_memory_id=ref_memory_id, - ) - for p in profiles - ] - - new_keys = {n.when_to_use for n in new_nodes} - original_count = len(nodes) - nodes = [n for n in nodes if n.when_to_use not in new_keys] - if len(nodes) < original_count: - logger.info(f"Removed {original_count - len(nodes)} duplicate profile(s) with matching keys") - - nodes.extend(new_nodes) - self._save_nodes(nodes) - logger.info(f"Batch added {len(new_nodes)} profiles") - return new_nodes - - def update_sync( - self, - profile_id: str, - message_time: str, - profile_key: str, - profile_value: str, - ) -> MemoryNode | None: - """Update fields for ``profile_id``; return ``None`` if that id is missing.""" - nodes = self._load_nodes() - target_node = None - for node in nodes: - if node.memory_id == profile_id: - node.when_to_use = profile_key - node.content = profile_value - node.message_time = message_time - target_node = node - break - - if target_node is None: - logger.warning(f"Profile {profile_id} not found") - return None - - self._save_nodes(nodes, apply_limits=False) - logger.info(f"Updated profile {profile_id}: {profile_key}={profile_value}") - return target_node - - def search_sync(self, query: str | list[str], limit: int = 5) -> list[MemoryNode]: - """Simple substring/token match over key and content, best matches first.""" - queries = [query] if isinstance(query, str) else query - query_terms = [q.strip().lower() for q in queries if q and q.strip()] - if not query_terms: - return [] - - scored_nodes = [] - for node in self.get_all_sync(): - profile_key = str(node.metadata.get("profile_key", node.when_to_use)).lower() - haystack = f"{profile_key}: {node.content}".lower() - score = 0 - for term in query_terms: - if term in haystack: - score += len(term) + 10 - else: - token_hits = sum(1 for token in term.split() if token and token in haystack) - score += token_hits - - if score > 0: - node.score = float(score) - scored_nodes.append(node) - - scored_nodes.sort(key=lambda n: (n.score, n.message_time), reverse=True) - return scored_nodes[:limit] - - async def get_all(self) -> list[MemoryNode]: - return self.get_all_sync() - - async def get_by(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None: - return self.get_by_sync(profile_id=profile_id, profile_key=profile_key) - - async def delete(self, profile_id: str | list[str]) -> bool | int: - return self.delete_sync(profile_id) - - async def delete_all(self) -> int: - return self.delete_all_sync() - - async def add( - self, - message_time: str, - profile_key: str, - profile_value: str, - ref_memory_id: str = "", - ) -> MemoryNode: - return self.add_sync(message_time, profile_key, profile_value, ref_memory_id) - - async def add_batch(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]: - return self.add_batch_sync(profiles, ref_memory_id) - - async def update( - self, - profile_id: str, - message_time: str, - profile_key: str, - profile_value: str, - ) -> MemoryNode | None: - return self.update_sync(profile_id, message_time, profile_key, profile_value) - - async def search(self, query: str | list[str], limit: int = 5) -> list[MemoryNode]: - return self.search_sync(query, limit) diff --git a/reme/memory/vector_tools/profiles/profile_backend.py b/reme/memory/vector_tools/profiles/profile_backend.py deleted file mode 100644 index 8ffcb9bc..00000000 --- a/reme/memory/vector_tools/profiles/profile_backend.py +++ /dev/null @@ -1,51 +0,0 @@ -"""Profile backend abstractions.""" - -from abc import ABC, abstractmethod - -from ....core.schema import MemoryNode - - -class BaseProfileBackend(ABC): - """Abstract interface for profile storage backends.""" - - def __init__(self, memory_target: str, max_capacity: int = 50): - self.memory_target = memory_target - self.max_capacity = max_capacity - - @abstractmethod - async def get_all(self) -> list[MemoryNode]: - """Return all profile rows for the current user.""" - - @abstractmethod - async def get_by(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None: - """Return one profile row by id or key.""" - - @abstractmethod - async def delete(self, profile_id: str | list[str]) -> bool | int: - """Delete one or more profile rows.""" - - @abstractmethod - async def delete_all(self) -> int: - """Delete all profile rows for the current user.""" - - @abstractmethod - async def add(self, message_time: str, profile_key: str, profile_value: str, ref_memory_id: str = "") -> MemoryNode: - """Add a single profile row.""" - - @abstractmethod - async def add_batch(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]: - """Add multiple profile rows.""" - - @abstractmethod - async def update( - self, - profile_id: str, - message_time: str, - profile_key: str, - profile_value: str, - ) -> MemoryNode | None: - """Update one profile row.""" - - @abstractmethod - async def search(self, query: str | list[str], limit: int = 5) -> list[MemoryNode]: - """Search profile rows relevant to the query.""" diff --git a/reme/memory/vector_tools/profiles/profile_handler.py b/reme/memory/vector_tools/profiles/profile_handler.py deleted file mode 100644 index 145caa6e..00000000 --- a/reme/memory/vector_tools/profiles/profile_handler.py +++ /dev/null @@ -1,195 +0,0 @@ -"""Profile handler facade for filesystem and vector backends.""" - -# pylint: disable=missing-function-docstring - -import asyncio -from pathlib import Path - -from loguru import logger - -from .file_profile_backend import FileProfileBackend -from .profile_backend import BaseProfileBackend -from .vector_profile_backend import VectorProfileBackend -from ....core import ServiceContext -from ....core.schema import MemoryNode - - -class ProfileHandler: - """User profile facade with pluggable storage backends.""" - - def __init__( - self, - memory_target: str, - profile_path: str | Path | None = None, - service_context: ServiceContext | None = None, - profile_backend: str = "filesystem", - profile_store_name: str = "profile", - max_capacity: int = 50, - ): - self.memory_target = memory_target - self.profile_backend = profile_backend - self.profile_store_name = profile_store_name - self.max_capacity = max_capacity - self.cache_key = self.memory_target.replace(" ", "_").lower() - self.backend = self._build_backend( - profile_path=profile_path, - service_context=service_context, - ) - - def _build_backend( - self, - profile_path: str | Path | None, - service_context: ServiceContext | None, - ) -> BaseProfileBackend: - if self.profile_backend == "filesystem": - if profile_path is None: - raise ValueError("profile_path is required for filesystem profile backend") - return FileProfileBackend( - profile_path=profile_path, - memory_target=self.memory_target, - max_capacity=self.max_capacity, - ) - - if self.profile_backend == "vector": - if service_context is None: - raise ValueError("service_context is required for vector profile backend") - return VectorProfileBackend( - memory_target=self.memory_target, - service_context=service_context, - vector_store_name=self.profile_store_name, - max_capacity=self.max_capacity, - ) - - raise ValueError(f"Unsupported profile backend: {self.profile_backend}") - - @staticmethod - def _run_sync(coro): - try: - asyncio.get_running_loop() - except RuntimeError: - return asyncio.run(coro) - raise RuntimeError( - "Synchronous profile access is not available in an active event loop. Use async methods instead.", - ) - - async def adelete(self, profile_id: str | list[str]) -> bool | int: - return await self.backend.delete(profile_id) - - async def adelete_all(self) -> int: - return await self.backend.delete_all() - - async def aadd( - self, - message_time: str, - profile_key: str, - profile_value: str, - ref_memory_id: str = "", - ) -> MemoryNode: - return await self.backend.add(message_time, profile_key, profile_value, ref_memory_id) - - async def aadd_batch(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]: - return await self.backend.add_batch(profiles, ref_memory_id) - - async def aupdate( - self, - profile_id: str, - message_time: str, - profile_key: str, - profile_value: str, - ) -> MemoryNode | None: - return await self.backend.update(profile_id, message_time, profile_key, profile_value) - - async def aget_by(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None: - return await self.backend.get_by(profile_id=profile_id, profile_key=profile_key) - - async def aget_by_id(self, profile_id: str) -> MemoryNode | None: - return await self.aget_by(profile_id=profile_id) - - async def aget_by_key(self, profile_key: str) -> MemoryNode | None: - return await self.aget_by(profile_key=profile_key) - - async def aget_all(self) -> list[MemoryNode]: - return await self.backend.get_all() - - async def asearch(self, query: str | list[str], limit: int = 5) -> list[MemoryNode]: - return await self.backend.search(query=query, limit=limit) - - @staticmethod - def format_node(node: MemoryNode, add_profile_id: bool = False, add_history_id: bool = False) -> str: - """Render a profile ``MemoryNode`` as a single-line string for tools/logs.""" - parts = [] - profile_key = str(node.metadata.get("profile_key", node.when_to_use)) - - if add_profile_id: - parts.append(f"profile_id={node.memory_id}") - - if node.message_time: - parts.append(f"[{node.message_time}]") - - parts.append(f"{profile_key}: {node.content}") - - if add_history_id and node.ref_memory_id: - parts.append(f"history_id={node.ref_memory_id}") - - return " ".join(parts) - - async def aread_all(self, add_profile_id: bool = False, add_history_id: bool = False) -> str: - nodes = await self.aget_all() - formatted_profiles = [self.format_node(node, add_profile_id, add_history_id) for node in nodes] - logger.info(f"Read {len(formatted_profiles)} profiles from {self.cache_key}") - return "\n".join(formatted_profiles).strip() - - async def aretrieve( - self, - query: str | list[str], - limit: int = 5, - add_profile_id: bool = True, - add_history_id: bool = False, - ) -> tuple[list[MemoryNode], str]: - nodes = await self.asearch(query=query, limit=limit) - formatted_profiles = [self.format_node(node, add_profile_id, add_history_id) for node in nodes] - return nodes, "\n".join(formatted_profiles).strip() - - def delete(self, profile_id: str | list[str]) -> bool | int: - if isinstance(self.backend, FileProfileBackend): - return self.backend.delete_sync(profile_id) - return self._run_sync(self.adelete(profile_id)) - - def delete_all(self) -> int: - if isinstance(self.backend, FileProfileBackend): - return self.backend.delete_all_sync() - return self._run_sync(self.adelete_all()) - - def add(self, message_time: str, profile_key: str, profile_value: str, ref_memory_id: str = "") -> MemoryNode: - if isinstance(self.backend, FileProfileBackend): - return self.backend.add_sync(message_time, profile_key, profile_value, ref_memory_id) - return self._run_sync(self.aadd(message_time, profile_key, profile_value, ref_memory_id)) - - def add_batch(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]: - if isinstance(self.backend, FileProfileBackend): - return self.backend.add_batch_sync(profiles, ref_memory_id) - return self._run_sync(self.aadd_batch(profiles, ref_memory_id)) - - def update(self, profile_id: str, message_time: str, profile_key: str, profile_value: str) -> MemoryNode | None: - if isinstance(self.backend, FileProfileBackend): - return self.backend.update_sync(profile_id, message_time, profile_key, profile_value) - return self._run_sync(self.aupdate(profile_id, message_time, profile_key, profile_value)) - - def get_by(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None: - if isinstance(self.backend, FileProfileBackend): - return self.backend.get_by_sync(profile_id=profile_id, profile_key=profile_key) - return self._run_sync(self.aget_by(profile_id=profile_id, profile_key=profile_key)) - - def get_by_id(self, profile_id: str) -> MemoryNode | None: - return self._run_sync(self.aget_by_id(profile_id)) - - def get_by_key(self, profile_key: str) -> MemoryNode | None: - return self._run_sync(self.aget_by_key(profile_key)) - - def get_all(self) -> list[MemoryNode]: - if isinstance(self.backend, FileProfileBackend): - return self.backend.get_all_sync() - return self._run_sync(self.aget_all()) - - def read_all(self, add_profile_id: bool = False, add_history_id: bool = False) -> str: - return self._run_sync(self.aread_all(add_profile_id, add_history_id)) diff --git a/reme/memory/vector_tools/profiles/profile_vector_handler.py b/reme/memory/vector_tools/profiles/profile_vector_handler.py deleted file mode 100644 index f03e6a96..00000000 --- a/reme/memory/vector_tools/profiles/profile_vector_handler.py +++ /dev/null @@ -1,245 +0,0 @@ -"""Vector-backed handler for bounded user profiles.""" - -import hashlib - -from loguru import logger - -from ....core import ServiceContext -from ....core.enumeration import MemoryType -from ....core.schema import MemoryNode -from ....core.vector_store import BaseVectorStore - - -class ProfileVectorHandler: - """Manage profile rows stored in a dedicated vector collection.""" - - PROFILE_KIND = "profile" - - def __init__( - self, - memory_target: str, - service_context: ServiceContext, - vector_store_name: str = "profile", - max_capacity: int = 50, - ): - self.memory_target = memory_target - self.service_context = service_context - self.vector_store_name = vector_store_name - self.max_capacity = max_capacity - self.vector_store: BaseVectorStore = service_context.vector_stores[vector_store_name] - - @staticmethod - def build_retrieval_text(profile_key: str, profile_value: str) -> str: - """Build the text that will be embedded for semantic profile retrieval.""" - return f"{profile_key}: {profile_value}".strip(": ") - - def build_profile_id(self, profile_key: str) -> str: - """Build a stable id from user and key.""" - hash_obj = hashlib.sha256(f"{self.memory_target}\n{profile_key}".encode("utf-8")) - return hash_obj.hexdigest()[:16] - - def _base_filters(self) -> dict: - """Filters shared by all profile rows in the vector collection.""" - return { - "memory_type": MemoryType.IDENTITY.value, - "memory_target": self.memory_target, - "profile_kind": self.PROFILE_KIND, - } - - def _build_profile_node(self, profile: dict, ref_memory_id: str = "") -> MemoryNode: - """Turn a profile dict into a ``MemoryNode`` for upsert into the vector store.""" - profile_key = profile.get("profile_key", "").strip() - profile_value = profile.get("profile_value", "").strip() - message_time = profile.get("message_time", "") - ref_id = profile.get("ref_memory_id", ref_memory_id) - metadata = dict(profile.get("metadata", {})) - metadata.update( - { - "profile_key": profile_key, - "profile_kind": self.PROFILE_KIND, - "profile_backend": "vector", - }, - ) - return MemoryNode( - memory_id=self.build_profile_id(profile_key), - memory_type=MemoryType.IDENTITY, - memory_target=self.memory_target, - when_to_use=self.build_retrieval_text(profile_key, profile_value), - content=profile_value, - message_time=message_time, - ref_memory_id=ref_id, - metadata=metadata, - ) - - def _vector_profile_matches(self, memory_node: MemoryNode) -> bool: - """True if ``memory_node`` belongs to this handler's target and profile kind.""" - if memory_node.memory_target != self.memory_target: - return False - if memory_node.memory_type is not MemoryType.IDENTITY: - return False - if memory_node.metadata.get("profile_kind") != self.PROFILE_KIND: - return False - return True - - async def _get_by_profile_id(self, profile_id: str) -> MemoryNode | None: - """Load by vector id and validate filters.""" - try: - vector_node = await self.vector_store.get(profile_id) - except KeyError: - logger.warning(f"Profile {profile_id} not found in vector store") - return None - if vector_node is None: - logger.warning(f"Profile {profile_id} not found in vector store") - return None - memory_node = MemoryNode.from_vector_node(vector_node) - if not self._vector_profile_matches(memory_node): - return None - return memory_node - - async def _get_by_profile_key(self, profile_key: str) -> MemoryNode | None: - """Load the single row matching ``profile_key`` under base filters.""" - filters = {**self._base_filters(), "profile_key": profile_key} - vector_nodes = await self.vector_store.list(filters=filters, limit=1) - if not vector_nodes: - return None - return MemoryNode.from_vector_node(vector_nodes[0]) - - async def get_all(self) -> list[MemoryNode]: - """List every profile row for this memory target, sorted by store.""" - vector_nodes = await self.vector_store.list( - filters=self._base_filters(), - sort_key="message_time", - reverse=False, - ) - return [MemoryNode.from_vector_node(node) for node in vector_nodes] - - async def get_by(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None: - """Return one profile by stable id or by logical profile key.""" - if not profile_id and not profile_key: - raise ValueError("Must provide either profile_id or profile_key") - if profile_id: - return await self._get_by_profile_id(profile_id) - return await self._get_by_profile_key(profile_key or "") - - async def delete(self, profile_id: str | list[str]) -> bool | int: - """Delete one id, many ids, or report zero/false when nothing matched.""" - if isinstance(profile_id, list): - profile_ids = list(dict.fromkeys(pid for pid in profile_id if pid)) - if not profile_ids: - return 0 - existing_nodes = [] - for pid in profile_ids: - node = await self.get_by(profile_id=pid) - if node is not None: - existing_nodes.append(node) - if not existing_nodes: - return 0 - await self.vector_store.delete([node.memory_id for node in existing_nodes]) - return len(existing_nodes) - - existing_node = await self.get_by(profile_id=profile_id) - if existing_node is None: - return False - await self.vector_store.delete(existing_node.memory_id) - return True - - async def delete_all(self) -> int: - """Remove all profile vectors for this target; returns how many were deleted.""" - nodes = await self.get_all() - if not nodes: - return 0 - await self.vector_store.delete([node.memory_id for node in nodes]) - return len(nodes) - - async def add_batch(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]: - """Upsert many profiles at once (last dict wins per key), then enforce capacity.""" - if not profiles: - return [] - - deduped_profiles: dict[str, dict] = {} - for profile in profiles: - profile_key = profile.get("profile_key", "").strip() - if not profile_key: - continue - deduped_profiles[profile_key] = profile - - new_nodes = [ - self._build_profile_node(profile, ref_memory_id=ref_memory_id) for profile in deduped_profiles.values() - ] - if not new_nodes: - return [] - - await self.vector_store.delete([node.memory_id for node in new_nodes]) - await self.vector_store.insert([node.to_vector_node() for node in new_nodes]) - await self.enforce_capacity() - return new_nodes - - async def add(self, message_time: str, profile_key: str, profile_value: str, ref_memory_id: str = "") -> MemoryNode: - """Insert or replace a single profile row.""" - nodes = await self.add_batch( - [ - { - "message_time": message_time, - "profile_key": profile_key, - "profile_value": profile_value, - }, - ], - ref_memory_id=ref_memory_id, - ) - return nodes[0] - - async def update( - self, - profile_id: str, - message_time: str, - profile_key: str, - profile_value: str, - ) -> MemoryNode | None: - """Replace content and key for ``profile_id``; return ``None`` if missing.""" - existing_node = await self.get_by(profile_id=profile_id) - if existing_node is None: - return None - - new_node = self._build_profile_node( - { - "message_time": message_time, - "profile_key": profile_key, - "profile_value": profile_value, - "ref_memory_id": existing_node.ref_memory_id, - "metadata": existing_node.metadata, - }, - ) - - if existing_node.memory_id != new_node.memory_id: - await self.vector_store.delete(existing_node.memory_id) - else: - await self.vector_store.delete(new_node.memory_id) - - await self.vector_store.insert(new_node.to_vector_node()) - await self.enforce_capacity() - return new_node - - async def search(self, query: str | list[str], limit: int = 5) -> list[MemoryNode]: - """Semantic search with de-duplication across multiple query strings.""" - queries = [query] if isinstance(query, str) else query - seen_nodes: dict[str, MemoryNode] = {} - for item in queries: - if not item or not item.strip(): - continue - vector_nodes = await self.vector_store.search(item, limit=limit, filters=self._base_filters()) - for vector_node in vector_nodes: - memory_node = MemoryNode.from_vector_node(vector_node) - seen_nodes[memory_node.memory_id] = memory_node - nodes = list(seen_nodes.values()) - nodes.sort(key=lambda node: (node.score, node.message_time), reverse=True) - return nodes[:limit] - - async def enforce_capacity(self): - """Drop oldest rows when count exceeds ``max_capacity``.""" - nodes = await self.get_all() - overflow = len(nodes) - self.max_capacity - if overflow <= 0: - return - - to_delete = [node.memory_id for node in nodes[:overflow]] - await self.vector_store.delete(to_delete) diff --git a/reme/memory/vector_tools/profiles/read_all_profiles.py b/reme/memory/vector_tools/profiles/read_all_profiles.py deleted file mode 100644 index 50eab684..00000000 --- a/reme/memory/vector_tools/profiles/read_all_profiles.py +++ /dev/null @@ -1,54 +0,0 @@ -"""Read user profile tool.""" - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall - - -class ReadAllProfiles(BaseMemoryTool): - """Tool to read all user profiles""" - - def __init__(self, enable_memory_target: bool = False, **kwargs): - kwargs["enable_multiple"] = False - super().__init__(**kwargs) - self.enable_memory_target: bool = enable_memory_target - - def _build_tool_call(self) -> ToolCall: - """Build and return the tool call schema""" - properties = {} - required = [] - - if self.enable_memory_target: - properties["memory_target"] = { - "type": "string", - "description": "memory_target", - } - required.append("memory_target") - - return ToolCall( - **{ - "description": "Read all user profiles.", - "parameters": { - "type": "object", - "properties": properties, - "required": required, - }, - }, - ) - - async def execute(self): - if self.enable_memory_target: - target = self.context.get("memory_target") - else: - target = self.memory_target - - profile_handler = self.get_profile_handler(target) - profiles_str = await profile_handler.aread_all(add_profile_id=True) - if not profiles_str: - output = "No profiles found." - logger.info(output) - return output - - logger.info("Successfully read profiles") - return profiles_str diff --git a/reme/memory/vector_tools/profiles/retrieve_profile.py b/reme/memory/vector_tools/profiles/retrieve_profile.py deleted file mode 100644 index 38337bd8..00000000 --- a/reme/memory/vector_tools/profiles/retrieve_profile.py +++ /dev/null @@ -1,102 +0,0 @@ -"""Retrieve relevant profile rows.""" - -from loguru import logger - -from .profile_handler import ProfileHandler -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import MemoryNode, ToolCall - - -class RetrieveProfile(BaseMemoryTool): - """Tool to retrieve relevant profiles using the configured backend.""" - - def __init__(self, top_k: int = 5, enable_memory_target: bool = False, **kwargs): - super().__init__(**kwargs) - self.top_k = top_k - self.enable_memory_target = enable_memory_target - - def _build_query_parameters(self) -> dict: - properties = { - "query": { - "type": "string", - "description": "query", - }, - } - required = ["query"] - if self.enable_memory_target: - properties["memory_target"] = { - "type": "string", - "description": "memory_target", - } - required.append("memory_target") - return { - "type": "object", - "properties": properties, - "required": required, - } - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "Retrieve relevant user profiles using semantic matching.", - "parameters": self._build_query_parameters(), - }, - ) - - def _build_multiple_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "Retrieve relevant user profiles using semantic matching.", - "parameters": { - "type": "object", - "properties": { - "query_items": { - "type": "array", - "description": "List of query items.", - "items": self._build_query_parameters(), - }, - }, - "required": ["query_items"], - }, - }, - ) - - async def execute(self): - if self.enable_multiple: - query_items = self.context.get("query_items", []) - else: - query_items = [self.context] - - queries_by_target: dict[str, list[str]] = {} - for item in query_items: - target = item["memory_target"] if self.enable_memory_target else self.memory_target - queries_by_target.setdefault(target, []).append(item["query"]) - - profile_nodes: list[MemoryNode] = [] - for target, queries in queries_by_target.items(): - profile_handler = self.get_profile_handler(target) - nodes, _ = await profile_handler.aretrieve( - query=queries, - limit=self.top_k, - add_profile_id=True, - add_history_id=True, - ) - profile_nodes.extend(nodes) - - seen_ids = {node.memory_id: node for node in self.retrieved_nodes if node.memory_id} - new_nodes = [] - for node in profile_nodes: - if node.memory_id not in seen_ids: - seen_ids[node.memory_id] = node - new_nodes.append(node) - self.retrieved_nodes.extend(new_nodes) - - if not new_nodes: - output = "No new profiles found." - else: - output = "\n".join( - [ProfileHandler.format_node(node, add_profile_id=True, add_history_id=True) for node in new_nodes], - ) - - logger.info(f"Retrieved {len(profile_nodes)} profiles, {len(new_nodes)} new after deduplication") - return output diff --git a/reme/memory/vector_tools/profiles/update_profile.py b/reme/memory/vector_tools/profiles/update_profile.py deleted file mode 100644 index 04c932bf..00000000 --- a/reme/memory/vector_tools/profiles/update_profile.py +++ /dev/null @@ -1,122 +0,0 @@ -"""Update user profile tool.""" - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall - - -class UpdateProfile(BaseMemoryTool): - """Tool to update user profile by adding or removing profile entries""" - - def __init__(self, enable_memory_target: bool = False, **kwargs): - kwargs["enable_multiple"] = True - super().__init__(**kwargs) - self.enable_memory_target: bool = enable_memory_target - - def _build_multiple_tool_call(self) -> ToolCall: - """Build and return the multiple tool call schema""" - profile_properties = { - "message_time": { - "type": "string", - "description": "Message time, e.g. '2020-01-01 00:00:00'", - }, - "profile_key": { - "type": "string", - "description": "Profile key or category, e.g. 'name'", - }, - "profile_value": { - "type": "string", - "description": "Profile value or content, e.g. 'John Smith'", - }, - } - profile_required = ["message_time", "profile_key", "profile_value"] - - if self.enable_memory_target: - profile_properties["memory_target"] = { - "type": "string", - "description": "memory_target", - } - profile_required.append("memory_target") - - return ToolCall( - **{ - "description": "update user profile by removing and adding profile entries.", - "parameters": { - "type": "object", - "properties": { - "profile_ids_to_delete": { - "type": "array", - "description": "List of profile IDs to delete", - "items": { - "type": "string", - }, - }, - "profiles_to_add": { - "type": "array", - "description": "List of profiles to add", - "items": { - "type": "object", - "properties": profile_properties, - "required": profile_required, - }, - }, - }, - "required": ["profile_ids_to_delete", "profiles_to_add"], - }, - }, - ) - - async def execute(self): - # Get parameters - profile_ids_to_delete = self.context.get("profile_ids_to_delete", []) - profile_ids_to_delete = sorted({pid for pid in profile_ids_to_delete if pid}) - profiles_to_add = self.context.get("profiles_to_add", []) - - if not profile_ids_to_delete and not profiles_to_add: - return "No profiles to remove or add, operation completed." - - removed_count = 0 - added_count = 0 - - # Delete profiles (using self.memory_target) - if profile_ids_to_delete: - profile_handler = self.get_profile_handler(self.memory_target) - removed_count = await profile_handler.adelete(profile_ids_to_delete) - - # Add new profiles - if profiles_to_add: - if self.enable_memory_target: - # Group profiles by memory_target - from collections import defaultdict - - profiles_by_target = defaultdict(list) - for profile in profiles_to_add: - target = profile.get("memory_target", self.memory_target) - profiles_by_target[target].append(profile) - - # Add profiles for each target - for target, target_profiles in profiles_by_target.items(): - profile_handler = self.get_profile_handler(target) - new_nodes = await profile_handler.aadd_batch( - profiles=target_profiles, - ref_memory_id=self.history_id, - ) - self.memory_nodes.extend(new_nodes) - added_count += len(new_nodes) - else: - # Use self.memory_target for all profiles - profile_handler = self.get_profile_handler(self.memory_target) - new_nodes = await profile_handler.aadd_batch(profiles=profiles_to_add, ref_memory_id=self.history_id) - self.memory_nodes.extend(new_nodes) - added_count = len(new_nodes) - - # Build output message - operations = [] - if removed_count > 0: - operations.append(f"removed {removed_count} old profiles.") - if added_count > 0: - operations.append(f"added {added_count} new profiles.") - operations.append("Operation completed.") - logger.info("\n".join(operations)) - return "\n".join(operations) diff --git a/reme/memory/vector_tools/profiles/update_profiles_v1.py b/reme/memory/vector_tools/profiles/update_profiles_v1.py deleted file mode 100644 index 128188fb..00000000 --- a/reme/memory/vector_tools/profiles/update_profiles_v1.py +++ /dev/null @@ -1,175 +0,0 @@ -"""Update user profile tool.""" - -from loguru import logger - -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall - - -class UpdateProfilesV1(BaseMemoryTool): - """Tool to update user profile by adding or removing profile entries""" - - def __init__(self, name="update_profiles", enable_memory_target: bool = False, **kwargs): - kwargs["enable_multiple"] = True - super().__init__(name=name, **kwargs) - self.enable_memory_target: bool = enable_memory_target - - def _build_profile_parameters(self, include_profile_id: bool = False) -> dict: - """Build the profile parameters schema based on enabled features.""" - properties = {} - required = [] - - if include_profile_id: - properties["profile_id"] = { - "type": "string", - "description": "ID of the profile to update", - } - required.append("profile_id") - - properties.update( - { - "message_time": { - "type": "string", - "description": "Message time, e.g. '2020-01-01 00:00:00'", - }, - "profile_key": { - "type": "string", - "description": "Profile key or category, e.g. 'name'", - }, - "profile_value": { - "type": "string", - "description": "Profile value or content, e.g. 'John Smith'", - }, - }, - ) - required.extend(["message_time", "profile_key", "profile_value"]) - - if self.enable_memory_target: - properties["memory_target"] = { - "type": "string", - "description": "memory_target", - } - required.append("memory_target") - - return { - "type": "object", - "properties": properties, - "required": required, - } - - def _build_multiple_tool_call(self) -> ToolCall: - """Build and return the multiple tool call schema""" - return ToolCall( - **{ - "description": "Update existing profiles and add new profiles.", - "parameters": { - "type": "object", - "properties": { - "profiles_to_update": { - "type": "array", - "description": "List of profiles to update", - "items": self._build_profile_parameters(include_profile_id=True), - }, - "profiles_to_add": { - "type": "array", - "description": "List of profiles to add", - "items": self._build_profile_parameters(include_profile_id=False), - }, - }, - "required": ["profiles_to_update", "profiles_to_add"], - }, - }, - ) - - async def execute(self): - # Get parameters - profiles_to_update = self.context.get("profiles_to_update", []) - profiles_to_add = self.context.get("profiles_to_add", []) - - if not profiles_to_update and not profiles_to_add: - return "No profiles to update or add, operation completed." - - # Step 1: Collect and delete all old profiles that need to be updated - if profiles_to_update: - # Group deletion IDs by memory_target if enabled - if self.enable_memory_target: - delete_by_target = {} - for profile in profiles_to_update: - target = profile.get("memory_target", self.memory_target) - profile_id = profile.get("profile_id") - if profile_id: - if target not in delete_by_target: - delete_by_target[target] = [] - delete_by_target[target].append(profile_id) - else: - delete_by_target = { - self.memory_target: [ - profile.get("profile_id") for profile in profiles_to_update if profile.get("profile_id") - ], - } - - # Delete old profiles for each target - for target, profile_ids in delete_by_target.items(): - if profile_ids: - profile_ids = sorted(set(profile_ids)) # Remove duplicates and sort - profile_handler = self.get_profile_handler(target) - await profile_handler.adelete(profile_ids) - - # Step 2: Prepare all profiles to add (both updated and new) - all_profiles_to_add = [] - - # Add profiles from updates - if profiles_to_update: - for profile in profiles_to_update: - target = ( - profile.get( - "memory_target", - self.memory_target, - ) - if self.enable_memory_target - else self.memory_target - ) - all_profiles_to_add.append((target, profile)) - - # Add new profiles - if profiles_to_add: - for profile in profiles_to_add: - target = ( - profile.get( - "memory_target", - self.memory_target, - ) - if self.enable_memory_target - else self.memory_target - ) - all_profiles_to_add.append((target, profile)) - - # Step 3: Group all profiles by target and add them in batch - from collections import defaultdict - - profiles_by_target = defaultdict(list) - for target, profile in all_profiles_to_add: - profiles_by_target[target].append(profile) - - # Process each target and add profiles - all_memory_nodes = [] - updated_count = len(profiles_to_update) - added_count = len(profiles_to_add) - - for target, target_profiles in profiles_by_target.items(): - profile_handler = self.get_profile_handler(target) - new_nodes = await profile_handler.aadd_batch(profiles=target_profiles, ref_memory_id=self.history_id) - all_memory_nodes.extend(new_nodes) - - # Extend memory_nodes for tracking - self.memory_nodes.extend(all_memory_nodes) - - # Build output message - operations = [] - if updated_count > 0: - operations.append(f"updated {updated_count} profiles.") - if added_count > 0: - operations.append(f"added {added_count} new profiles.") - operations.append("Operation completed.") - logger.info("\n".join(operations)) - return "\n".join(operations) diff --git a/reme/memory/vector_tools/profiles/vector_profile_backend.py b/reme/memory/vector_tools/profiles/vector_profile_backend.py deleted file mode 100644 index 3147cc92..00000000 --- a/reme/memory/vector_tools/profiles/vector_profile_backend.py +++ /dev/null @@ -1,55 +0,0 @@ -"""Vector-backed profile storage.""" - -from .profile_backend import BaseProfileBackend -from .profile_vector_handler import ProfileVectorHandler -from ....core import ServiceContext -from ....core.schema import MemoryNode - - -class VectorProfileBackend(BaseProfileBackend): - """Persist user profiles in a dedicated vector store.""" - - def __init__( - self, - memory_target: str, - service_context: ServiceContext, - vector_store_name: str = "profile", - max_capacity: int = 50, - ): - super().__init__(memory_target=memory_target, max_capacity=max_capacity) - self.handler = ProfileVectorHandler( - memory_target=memory_target, - service_context=service_context, - vector_store_name=vector_store_name, - max_capacity=max_capacity, - ) - - async def get_all(self) -> list[MemoryNode]: - return await self.handler.get_all() - - async def get_by(self, *, profile_id: str | None = None, profile_key: str | None = None) -> MemoryNode | None: - return await self.handler.get_by(profile_id=profile_id, profile_key=profile_key) - - async def delete(self, profile_id: str | list[str]) -> bool | int: - return await self.handler.delete(profile_id) - - async def delete_all(self) -> int: - return await self.handler.delete_all() - - async def add(self, message_time: str, profile_key: str, profile_value: str, ref_memory_id: str = "") -> MemoryNode: - return await self.handler.add(message_time, profile_key, profile_value, ref_memory_id) - - async def add_batch(self, profiles: list[dict], ref_memory_id: str = "") -> list[MemoryNode]: - return await self.handler.add_batch(profiles, ref_memory_id) - - async def update( - self, - profile_id: str, - message_time: str, - profile_key: str, - profile_value: str, - ) -> MemoryNode | None: - return await self.handler.update(profile_id, message_time, profile_key, profile_value) - - async def search(self, query: str | list[str], limit: int = 5) -> list[MemoryNode]: - return await self.handler.search(query, limit) diff --git a/reme/memory/vector_tools/record/__init__.py b/reme/memory/vector_tools/record/__init__.py deleted file mode 100644 index e69de29b..00000000 diff --git a/reme/memory/vector_tools/record/add_and_retrieve_similar_memory.py b/reme/memory/vector_tools/record/add_and_retrieve_similar_memory.py deleted file mode 100644 index 00aa0383..00000000 --- a/reme/memory/vector_tools/record/add_and_retrieve_similar_memory.py +++ /dev/null @@ -1,127 +0,0 @@ -"""Add draft memory and retrieve similar memories from vector store""" - -from loguru import logger - -from .memory_handler import MemoryHandler -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall, MemoryNode -from ....core.utils import deduplicate_memories - - -class AddAndRetrieveSimilarMemory(BaseMemoryTool): - """Tool to add draft memory and retrieve similar memories""" - - def __init__( - self, - top_k: int = 20, - enable_memory_target: bool = False, - enable_when_to_use: bool = False, - **kwargs, - ): - super().__init__(**kwargs) - self.top_k: int = top_k - self.enable_memory_target: bool = enable_memory_target - self.enable_when_to_use: bool = enable_when_to_use - - def _build_query_parameters(self) -> dict: - """Build the query parameters schema""" - properties = { - "message_time": { - "type": "string", - "description": "message time, e.g. '2020-01-01 00:00:00'", - }, - "memory_content": { - "type": "string", - "description": "content of the memory.", - }, - } - required = ["message_time", "memory_content"] - - if self.enable_when_to_use: - properties["when_to_use"] = { - "type": "string", - "description": "description of when to use this memory.", - } - required.append("when_to_use") - - if self.enable_memory_target: - properties["memory_target"] = { - "type": "string", - "description": "target memory type for this memory.", - } - required.append("memory_target") - - return { - "type": "object", - "properties": properties, - "required": required, - } - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "Add memory and retrieve similar memories from the vector store.", - "parameters": self._build_query_parameters(), - }, - ) - - def _build_multiple_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "Add memory and retrieve similar memories from the vector store.", - "parameters": { - "type": "object", - "properties": { - "items": { - "type": "array", - "description": "items", - "items": self._build_query_parameters(), - }, - }, - "required": ["items"], - }, - }, - ) - - async def execute(self): - if self.enable_multiple: - draft_items = self.context.get("items", []) - else: - draft_items = [self.context] - - queries_by_target: dict[str, list[dict]] = {} - for item in draft_items: - if self.enable_memory_target: - target = item["memory_target"] - else: - target = self.memory_target - if target not in queries_by_target: - queries_by_target[target] = [] - - queries_by_target[target].append( - { - "query": item["memory_content"], - "limit": self.top_k, - "filters": {}, - }, - ) - - # Execute batch searches for each target - memory_nodes: list[MemoryNode] = [] - for target, searches in queries_by_target.items(): - handler = MemoryHandler(target, self.service_context) - nodes = await handler.batch_search(searches) - memory_nodes.extend(nodes) - - memory_nodes = deduplicate_memories(memory_nodes) - retrieved_ids = {n.memory_id for n in self.retrieved_nodes if n.memory_id} - new_nodes = [n for n in memory_nodes if n.memory_id not in retrieved_ids] - self.retrieved_nodes.extend(new_nodes) - - if not new_nodes: - output = "No similar memories found." - else: - output = "\n".join([n.format(ref_memory_id_key="history_id") for n in new_nodes]) - - logger.info(f"Retrieved {len(memory_nodes)} similar memories, {len(new_nodes)} new after deduplication") - return output diff --git a/reme/memory/vector_tools/record/add_draft_and_retrieve_similar_memory.py b/reme/memory/vector_tools/record/add_draft_and_retrieve_similar_memory.py deleted file mode 100644 index 040c475a..00000000 --- a/reme/memory/vector_tools/record/add_draft_and_retrieve_similar_memory.py +++ /dev/null @@ -1,127 +0,0 @@ -"""Add draft memory and retrieve similar memories from vector store""" - -from loguru import logger - -from .memory_handler import MemoryHandler -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall, MemoryNode -from ....core.utils import deduplicate_memories - - -class AddDraftAndRetrieveSimilarMemory(BaseMemoryTool): - """Tool to add draft memory and retrieve similar memories""" - - def __init__( - self, - top_k: int = 20, - enable_memory_target: bool = False, - enable_when_to_use: bool = False, - **kwargs, - ): - super().__init__(**kwargs) - self.top_k: int = top_k - self.enable_memory_target: bool = enable_memory_target - self.enable_when_to_use: bool = enable_when_to_use - - def _build_query_parameters(self) -> dict: - """Build the query parameters schema""" - properties = { - "message_time": { - "type": "string", - "description": "message time, e.g. '2020-01-01 00:00:00'", - }, - "memory_content": { - "type": "string", - "description": "content of the memory.", - }, - } - required = ["message_time", "memory_content"] - - if self.enable_when_to_use: - properties["when_to_use"] = { - "type": "string", - "description": "description of when to use this memory.", - } - required.append("when_to_use") - - if self.enable_memory_target: - properties["memory_target"] = { - "type": "string", - "description": "target memory type for this memory.", - } - required.append("memory_target") - - return { - "type": "object", - "properties": properties, - "required": required, - } - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "Add draft memory and retrieve similar memories from the vector store.", - "parameters": self._build_query_parameters(), - }, - ) - - def _build_multiple_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "Add draft memory and retrieve similar memories from the vector store.", - "parameters": { - "type": "object", - "properties": { - "draft_items": { - "type": "array", - "description": "draft_items", - "items": self._build_query_parameters(), - }, - }, - "required": ["draft_items"], - }, - }, - ) - - async def execute(self): - if self.enable_multiple: - draft_items = self.context.get("draft_items", []) - else: - draft_items = [self.context] - - queries_by_target: dict[str, list[dict]] = {} - for item in draft_items: - if self.enable_memory_target: - target = item["memory_target"] - else: - target = self.memory_target - if target not in queries_by_target: - queries_by_target[target] = [] - - queries_by_target[target].append( - { - "query": item["memory_content"], - "limit": self.top_k, - "filters": {}, - }, - ) - - # Execute batch searches for each target - memory_nodes: list[MemoryNode] = [] - for target, searches in queries_by_target.items(): - handler = MemoryHandler(target, self.service_context) - nodes = await handler.batch_search(searches) - memory_nodes.extend(nodes) - - memory_nodes = deduplicate_memories(memory_nodes) - retrieved_ids = {n.memory_id for n in self.retrieved_nodes if n.memory_id} - new_nodes = [n for n in memory_nodes if n.memory_id not in retrieved_ids] - self.retrieved_nodes.extend(new_nodes) - - if not new_nodes: - output = "No similar memories found." - else: - output = "\n".join([n.format(ref_memory_id_key="history_id") for n in new_nodes]) - - logger.info(f"Retrieved {len(memory_nodes)} similar memories, {len(new_nodes)} new after deduplication") - return output diff --git a/reme/memory/vector_tools/record/add_memory.py b/reme/memory/vector_tools/record/add_memory.py deleted file mode 100644 index 3a049fc9..00000000 --- a/reme/memory/vector_tools/record/add_memory.py +++ /dev/null @@ -1,137 +0,0 @@ -"""Add memory to vector store""" - -from loguru import logger - -from .memory_handler import MemoryHandler -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall - - -class AddMemory(BaseMemoryTool): - """Tool to add memories to vector store""" - - def __init__( - self, - enable_memory_target: bool = False, - enable_when_to_use: bool = False, - **kwargs, - ): - super().__init__(**kwargs) - self.enable_memory_target: bool = enable_memory_target - self.enable_when_to_use: bool = enable_when_to_use - - def _build_memory_parameters(self) -> dict: - """Build the memory parameters schema based on enabled features.""" - properties = { - "message_time": { - "type": "string", - "description": "message time, e.g. '2020-01-01 00:00:00'", - }, - "memory_content": { - "type": "string", - "description": "content of the memory.", - }, - } - required = ["message_time", "memory_content"] - - if self.enable_when_to_use: - properties["when_to_use"] = { - "type": "string", - "description": "description of when to use this memory.", - } - required.append("when_to_use") - - if self.enable_memory_target: - properties["memory_target"] = { - "type": "string", - "description": "target memory type for this memory.", - } - required.append("memory_target") - - return { - "type": "object", - "properties": properties, - "required": required, - } - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "add a memory to vector store for future retrieval.", - "parameters": self._build_memory_parameters(), - }, - ) - - def _build_multiple_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "add multiple memories to vector store for future retrieval.", - "parameters": { - "type": "object", - "properties": { - "memories": { - "type": "array", - "description": "list of memories to store.", - "items": self._build_memory_parameters(), - }, - }, - "required": ["memories"], - }, - }, - ) - - async def execute(self): - if self.enable_multiple: - memories = self.context.get("memories", []) - else: - memories = [self.context] - - # Group memories by memory_target if enabled - if self.enable_memory_target: - memories_by_target = {} - for mem in memories: - target = mem["memory_target"] - if target not in memories_by_target: - memories_by_target[target] = [] - memories_by_target[target].append(mem) - else: - memories_by_target = {self.memory_target: memories} - - # Process each memory_target group - all_memory_nodes = [] - for target, target_memories in memories_by_target.items(): - # Parse and prepare memory data - memory_dicts = [] - for mem in target_memories: - memory_content = mem.get("memory_content", "") - message_time = mem.get("message_time", "") - when_to_use = mem.get("when_to_use", "") if self.enable_when_to_use else "" - metadata = {} - try: - metadata["time_int"] = int(message_time.split(" ")[0].replace("-", "")) - except Exception: - logger.warning(f"Invalid message time format: {message_time}") - - memory_dicts.append( - { - "content": memory_content, - "when_to_use": when_to_use, - "message_time": message_time, - "ref_memory_id": self.history_id, - "author": self.author, - "metadata": metadata, - }, - ) - - if memory_dicts: - handler = MemoryHandler(target, self.service_context) - memory_nodes = await handler.add_batch(memory_dicts) - all_memory_nodes.extend(memory_nodes) - - if not all_memory_nodes: - return "No valid memories provided." - - self.memory_nodes.extend(all_memory_nodes) - output = f"Successfully added {len(all_memory_nodes)} memories." - logger.info(output) - return output diff --git a/reme/memory/vector_tools/record/delete_memory.py b/reme/memory/vector_tools/record/delete_memory.py deleted file mode 100644 index da8094fe..00000000 --- a/reme/memory/vector_tools/record/delete_memory.py +++ /dev/null @@ -1,59 +0,0 @@ -"""Delete memory from vector store""" - -from loguru import logger - -from .memory_handler import MemoryHandler -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall - - -class DeleteMemory(BaseMemoryTool): - """Tool to delete memories from vector store""" - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "delete a memory from vector store using its unique ID.", - "parameters": { - "type": "object", - "properties": { - "memory_id": { - "type": "string", - "description": "memory_id of the memory to delete.", - }, - }, - "required": ["memory_id"], - }, - }, - ) - - def _build_multiple_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "delete multiple memories from vector store using their unique IDs.", - "parameters": { - "type": "object", - "properties": { - "memory_ids": { - "type": "array", - "description": "memory_ids of memories to delete.", - "items": {"type": "string"}, - }, - }, - "required": ["memory_ids"], - }, - }, - ) - - async def execute(self): - memory_ids = self.context.get("memory_ids") or [] - if not memory_ids: - memory_ids = [self.context.get("memory_id", "")] - - handler = MemoryHandler(self.memory_target, self.service_context) - await handler.delete(memory_ids) - self.memory_nodes.extend(memory_ids) - - output = f"Successfully deleted {len(memory_ids)} memories." - logger.info(output) - return output diff --git a/reme/memory/vector_tools/record/memory_handler.py b/reme/memory/vector_tools/record/memory_handler.py deleted file mode 100644 index 7997ac86..00000000 --- a/reme/memory/vector_tools/record/memory_handler.py +++ /dev/null @@ -1,287 +0,0 @@ -"""Memory handler""" - -import numpy as np -from loguru import logger - -from ....core import ServiceContext -from ....core.enumeration import MemoryType -from ....core.schema import MemoryNode -from ....core.utils.common_utils import batch_cosine_similarity -from ....core.vector_store import BaseVectorStore - - -class MemoryHandler: - """Handler for managing memory nodes in the vector store.""" - - def __init__(self, memory_target: str, service_context: ServiceContext): - self.memory_target: str = memory_target - self.memory_type: MemoryType | None = service_context.memory_target_type_mapping.get(memory_target, None) - self.vector_store: BaseVectorStore = service_context.vector_stores["default"] - - async def add_batch(self, memories: list[dict]) -> list[MemoryNode]: - """Add multiple memory nodes and return their memory_ids.""" - # First, delete existing memory nodes if memory_ids are provided - memory_ids_to_delete = [mem.get("memory_id") for mem in memories if mem.get("memory_id")] - if memory_ids_to_delete: - await self.vector_store.delete(memory_ids_to_delete) - - # Create MemoryNode objects - memory_nodes = [ - MemoryNode( - memory_type=self.memory_type, - memory_target=self.memory_target, - content=mem.get("content", ""), - when_to_use=mem.get("when_to_use", ""), - message_time=mem.get("message_time", ""), - ref_memory_id=mem.get("ref_memory_id", ""), - author=mem.get("author", ""), - score=mem.get("score", 0.0), - metadata=mem.get("metadata", {}), - ) - for mem in memories - ] - - # Deduplicate memory_nodes by content (keep last occurrence) - memory_dict = {node.content: node for node in memory_nodes} - memory_nodes = list(memory_dict.values()) - - # Convert to VectorNodes and insert - vector_nodes = [node.to_vector_node() for node in memory_nodes] - await self.vector_store.insert(vector_nodes) - return memory_nodes - - async def add( - self, - content: str, - when_to_use: str = "", - message_time: str = "", - ref_memory_id: str = "", - author: str = "", - score: float = 0.0, - **kwargs, - ) -> MemoryNode: - """Add a single memory node and return its memory_id.""" - memory_dict = { - "content": content, - "when_to_use": when_to_use, - "message_time": message_time, - "ref_memory_id": ref_memory_id, - "author": author, - "score": score, - "metadata": kwargs, - } - memory_nodes = await self.add_batch([memory_dict]) - return memory_nodes[0] - - async def get(self, memory_ids: str | list[str]) -> MemoryNode | list[MemoryNode]: - """Get one or more memory nodes by their memory_ids.""" - # Call vector_store.get with the memory_ids - vector_nodes = await self.vector_store.get(memory_ids) - - # Convert VectorNode(s) to MemoryNode(s) - if isinstance(vector_nodes, list): - return [MemoryNode.from_vector_node(node) for node in vector_nodes] - else: - return MemoryNode.from_vector_node(vector_nodes) - - async def delete(self, memory_ids: str | list[str]): - """Delete multiple memory nodes by their memory_ids.""" - # Deduplicate if input is a list - if isinstance(memory_ids, list): - memory_ids = list(dict.fromkeys(memory_ids)) - await self.vector_store.delete(memory_ids) - - async def delete_all(self): - """Delete all memory nodes.""" - await self.vector_store.delete_all() - - async def update_batch(self, updates: list[dict]) -> list[MemoryNode]: - """Update multiple memory nodes with their memory_ids and new values using delete + add.""" - # Deduplicate updates by memory_id (keep last occurrence) - updates_dict = {upd["memory_id"]: upd for upd in updates} - updates = list(updates_dict.values()) - memory_ids = list(updates_dict.keys()) - - # Get existing nodes - vector_nodes = await self.vector_store.get(memory_ids) - if not isinstance(vector_nodes, list): - vector_nodes = [vector_nodes] - - # Update and convert back - updated_nodes: list[MemoryNode] = [] - for vector_node, update in zip(vector_nodes, updates): - memory_node = MemoryNode.from_vector_node(vector_node) - memory_node.memory_target = self.memory_target - memory_node.memory_type = self.memory_type - - if "content" in update: - memory_node.content = update["content"] - if "when_to_use" in update: - memory_node.when_to_use = update["when_to_use"] - if "message_time" in update: - memory_node.message_time = update["message_time"] - if "ref_memory_id" in update: - memory_node.ref_memory_id = update["ref_memory_id"] - if "author" in update: - memory_node.author = update["author"] - if "score" in update: - memory_node.score = update["score"] - if "metadata" in update: - memory_node.metadata.update(update["metadata"]) - updated_nodes.append(memory_node) - - # Delete old nodes first - await self.vector_store.delete(memory_ids) - - # Then add updated nodes - vector_nodes = [node.to_vector_node() for node in updated_nodes] - await self.vector_store.insert(vector_nodes) - - return updated_nodes - - async def update( - self, - memory_id: str, - content: str | None = None, - when_to_use: str | None = None, - message_time: str | None = None, - ref_memory_id: str | None = None, - author: str | None = None, - score: float | None = None, - **kwargs, - ) -> MemoryNode: - """Update a memory node's content, when_to_use, or other fields.""" - update_dict: dict = {"memory_id": memory_id} - if content is not None: - update_dict["content"] = content - if when_to_use is not None: - update_dict["when_to_use"] = when_to_use - if message_time is not None: - update_dict["message_time"] = message_time - if ref_memory_id is not None: - update_dict["ref_memory_id"] = ref_memory_id - if author is not None: - update_dict["author"] = author - if score is not None: - update_dict["score"] = score - if kwargs is not None: - update_dict["metadata"] = kwargs - - memory_nodes = await self.update_batch([update_dict]) - return memory_nodes[0] - - async def search( - self, - query: str | list[str], - limit: int = 5, - filters: dict | None = None, - **kwargs, - ) -> list[MemoryNode]: - """Search for similar memory nodes based on query text.""" - filters = filters or {} - filters["memory_type"] = self.memory_type.value - filters["memory_target"] = self.memory_target - - # Handle single query - if isinstance(query, str): - vector_nodes = await self.vector_store.search(query, limit=limit, filters=filters, **kwargs) - return [MemoryNode.from_vector_node(node) for node in vector_nodes] - - # Handle multiple queries: search each query with the same limit - seen_ids: dict[str, MemoryNode] = {} - - for q in query: - vector_nodes = await self.vector_store.search(q, limit=limit, filters=filters, **kwargs) - for vector_node in vector_nodes: - memory_node = MemoryNode.from_vector_node(vector_node) - if memory_node.memory_id not in seen_ids: - seen_ids[memory_node.memory_id] = memory_node - - return list(seen_ids.values()) - - async def batch_search(self, searches: list[dict], hybrid_threshold: float = None) -> list[MemoryNode]: - """Execute multiple search queries in batch and return deduplicated results.""" - if hybrid_threshold is not None: - # Extract query list from searches - query_list: list[str] = [search["query"] for search in searches] - - # Step 1: Get embeddings for all queries using the embedding model - # Shape: [query_size X emb_size] - embedding_model = self.vector_store.embedding_model - query_embeddings_list: list[list[float]] = await embedding_model.get_embeddings(query_list) - query_embeddings = np.array(query_embeddings_list) # Convert to numpy array - - # Step 2: Use self.search to get search results for each query and deduplicate - seen_ids: dict[str, MemoryNode] = {} - for search_params in searches: - search_result = await self.search(**search_params) - for memory_node in search_result: - if memory_node.memory_id not in seen_ids: - seen_ids[memory_node.memory_id] = memory_node - - # Step 3: Get deduplicated results - deduplicated_results = list(seen_ids.values()) - - # If no results, return empty list - if not deduplicated_results: - return [] - - # Step 4: Extract embeddings from results - # Shape: [result_size X emb_size] - result_embeddings_list = [node.vector for node in deduplicated_results if node.vector] - - # Filter out nodes without embeddings - results_with_embeddings = [node for node in deduplicated_results if node.vector] - - if not result_embeddings_list: - logger.warning("No results with embeddings found") - return deduplicated_results - - result_embeddings = np.array(result_embeddings_list) - - # Step 5: Compute cosine similarity matrix - # Shape: [query_size X result_size] - similarity_matrix = batch_cosine_similarity(query_embeddings, result_embeddings) - - # Step 6: Calculate average score for each result across all queries - # Shape: [result_size] - avg_scores = np.mean(similarity_matrix, axis=0) - - # Step 7: Filter results by hybrid_threshold and sort by average score - filtered_results = [] - for idx, node in enumerate(results_with_embeddings): - if avg_scores[idx] >= hybrid_threshold: - node.score = float(avg_scores[idx]) - filtered_results.append(node) - - # Sort by score in descending order - filtered_results.sort(key=lambda x: x.score, reverse=True) - - return filtered_results - - else: - # Original behavior: simple deduplication without hybrid scoring - seen_ids: dict[str, MemoryNode] = {} - - for search_params in searches: - search_result = await self.search(**search_params) - for memory_node in search_result: - if memory_node.memory_id not in seen_ids: - seen_ids[memory_node.memory_id] = memory_node - - return list(seen_ids.values()) - - async def list( - self, - filters: dict | None = None, - limit: int | None = None, - sort_key: str | None = None, - reverse: bool = True, - ) -> list[MemoryNode]: - """List memory nodes with optional filtering and sorting.""" - filters = filters or {} - filters["memory_type"] = self.memory_type.value - filters["memory_target"] = self.memory_target - - vector_nodes = await self.vector_store.list(filters=filters, limit=limit, sort_key=sort_key, reverse=reverse) - return [MemoryNode.from_vector_node(node) for node in vector_nodes] diff --git a/reme/memory/vector_tools/record/retrieve_memory.py b/reme/memory/vector_tools/record/retrieve_memory.py deleted file mode 100644 index 4fd30c9d..00000000 --- a/reme/memory/vector_tools/record/retrieve_memory.py +++ /dev/null @@ -1,140 +0,0 @@ -"""Retrieve memory from vector store""" - -from loguru import logger - -from .memory_handler import MemoryHandler -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall, MemoryNode -from ....core.utils import deduplicate_memories - - -class RetrieveMemory(BaseMemoryTool): - """Tool to retrieve memories using similarity search""" - - def __init__( - self, - top_k: int = 20, - enable_memory_target: bool = False, - enable_time_filter: bool = False, - hybrid_threshold: float | None = None, - **kwargs, - ): - super().__init__(**kwargs) - self.top_k: int = top_k - self.enable_memory_target: bool = enable_memory_target - self.enable_time_filter: bool = enable_time_filter - self.hybrid_threshold: float | None = hybrid_threshold - - def _build_query_parameters(self) -> dict: - """Build the query parameters schema based on enabled features.""" - properties = { - "query": { - "type": "string", - "description": "query", - }, - } - required = ["query"] - - if self.enable_time_filter: - properties["time_filter"] = { - "type": "string", - "description": "Optional time filter to narrow down search results by date. " - "Format: single date '20200101' for exact date match, " - "or date range '20200101,20200102' for inclusive range filtering.", - } - - if self.enable_memory_target: - properties["memory_target"] = { - "type": "string", - "description": "memory_target", - } - required.append("memory_target") - - return { - "type": "object", - "properties": properties, - "required": required, - } - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "Retrieve relevant memories from the vector store using semantic similarity search.", - "parameters": self._build_query_parameters(), - }, - ) - - def _build_multiple_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "Retrieve relevant memories from the vector store using semantic similarity search.", - "parameters": { - "type": "object", - "properties": { - "query_items": { - "type": "array", - "description": "List of query items.", - "items": self._build_query_parameters(), - }, - }, - "required": ["query_items"], - }, - }, - ) - - async def execute(self): - if self.enable_multiple: - query_items = self.context.get("query_items", []) - else: - query_items = [self.context] - - queries_by_target: dict[str, list[dict]] = {} - for item in query_items: - if self.enable_memory_target: - target = item["memory_target"] - else: - target = self.memory_target - if target not in queries_by_target: - queries_by_target[target] = [] - - filters = {} - time_filter = item.get("time_filter") - if time_filter: - time_filter = time_filter.strip() - if "," in time_filter: - start, end = time_filter.split(",") - filters = {"time_int": [int(start.strip()), int(end.strip())]} - else: - filters = {"time_int": [int(time_filter), int(time_filter)]} - - queries_by_target[target].append( - { - "query": item["query"], - "limit": self.top_k, - "filters": filters, - }, - ) - - # Execute batch searches for each target - memory_nodes: list[MemoryNode] = [] - for target, searches in queries_by_target.items(): - handler = MemoryHandler(target, self.service_context) - if self.hybrid_threshold is not None: - nodes = await handler.batch_search(searches, self.hybrid_threshold) - nodes = nodes[: self.top_k] - else: - nodes = await handler.batch_search(searches) - memory_nodes.extend(nodes) - - memory_nodes = deduplicate_memories(memory_nodes) - retrieved_ids = {n.memory_id for n in self.retrieved_nodes if n.memory_id} - new_nodes = [n for n in memory_nodes if n.memory_id not in retrieved_ids] - self.retrieved_nodes.extend(new_nodes) - - if not new_nodes: - output = "No new memories found." - else: - output = "\n".join([n.format(ref_memory_id_key="history_id") for n in new_nodes]) - - logger.info(f"Retrieved {len(memory_nodes)} memories, {len(new_nodes)} new after deduplication") - return output diff --git a/reme/memory/vector_tools/record/retrieve_recent_memory.py b/reme/memory/vector_tools/record/retrieve_recent_memory.py deleted file mode 100644 index fc94d23d..00000000 --- a/reme/memory/vector_tools/record/retrieve_recent_memory.py +++ /dev/null @@ -1,52 +0,0 @@ -"""Retrieve most recent memories from vector store""" - -from loguru import logger - -from .memory_handler import MemoryHandler -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall, MemoryNode -from ....core.utils import deduplicate_memories - - -class RetrieveRecentMemory(BaseMemoryTool): - """Tool to retrieve most recent memories sorted by time""" - - def __init__(self, top_k: int = 20, **kwargs): - kwargs["enable_multiple"] = False - super().__init__(**kwargs) - self.top_k: int = top_k - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "retrieve the most recent memories sorted by message time (newest first).", - "parameters": { - "type": "object", - "properties": {}, - "required": [], - }, - }, - ) - - async def execute(self): - handler = MemoryHandler(self.memory_target, self.service_context) - - memory_nodes: list[MemoryNode] = await handler.list( - limit=self.top_k, - sort_key="message_time", - reverse=True, - ) - memory_nodes = deduplicate_memories(memory_nodes) - - retrieved_ids = {n.memory_id for n in self.retrieved_nodes if n.memory_id} - new_nodes = [n for n in memory_nodes if n.memory_id not in retrieved_ids] - self.retrieved_nodes.extend(new_nodes) - self.memory_nodes.extend(new_nodes) - - if not new_nodes: - output = "No new memories found." - else: - output = "\n".join([n.format(ref_memory_id_key="history_id") for n in new_nodes]) - - logger.info(f"Retrieved {len(memory_nodes)} memories, {len(new_nodes)} new after deduplication") - return output diff --git a/reme/memory/vector_tools/record/update_memory.py b/reme/memory/vector_tools/record/update_memory.py deleted file mode 100644 index dbdf9ce6..00000000 --- a/reme/memory/vector_tools/record/update_memory.py +++ /dev/null @@ -1,141 +0,0 @@ -"""Update memory in vector store""" - -from loguru import logger - -from .memory_handler import MemoryHandler -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall - - -class UpdateMemory(BaseMemoryTool): - """Tool to update memories in vector store""" - - def __init__( - self, - enable_memory_target: bool = False, - enable_when_to_use: bool = False, - **kwargs, - ): - super().__init__(**kwargs) - self.enable_memory_target: bool = enable_memory_target - self.enable_when_to_use: bool = enable_when_to_use - - def _build_update_parameters(self) -> dict: - """Build the update parameters schema based on enabled features.""" - properties = { - "memory_id": { - "type": "string", - "description": "unique identifier of memory to update.", - }, - "message_time": { - "type": "string", - "description": "message time, e.g. '2020-01-01 00:00:00'", - }, - "memory_content": { - "type": "string", - "description": "new content of the memory.", - }, - } - required = ["memory_id", "message_time", "memory_content"] - - if self.enable_when_to_use: - properties["when_to_use"] = { - "type": "string", - "description": "description of when to use this memory.", - } - required.append("when_to_use") - - if self.enable_memory_target: - properties["memory_target"] = { - "type": "string", - "description": "target memory type for this memory.", - } - required.append("memory_target") - - return { - "type": "object", - "properties": properties, - "required": required, - } - - def _build_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "update a memory in vector store by replacing old memory with new content.", - "parameters": self._build_update_parameters(), - }, - ) - - def _build_multiple_tool_call(self) -> ToolCall: - return ToolCall( - **{ - "description": "update multiple memories in vector store by replacing old memories with new content.", - "parameters": { - "type": "object", - "properties": { - "memories": { - "type": "array", - "description": "list of memory update objects.", - "items": self._build_update_parameters(), - }, - }, - "required": ["memories"], - }, - }, - ) - - async def execute(self): - if self.enable_multiple: - memories = self.context.get("memories", []) - else: - memories = [self.context] - - # Group memories by memory_target if enabled - if self.enable_memory_target: - memories_by_target = {} - for mem in memories: - target = mem["memory_target"] - if target not in memories_by_target: - memories_by_target[target] = [] - memories_by_target[target].append(mem) - else: - memories_by_target = {self.memory_target: memories} - - # Process each memory_target group - all_memory_nodes = [] - for target, target_memories in memories_by_target.items(): - # Parse and prepare update data - update_dicts = [] - for mem in target_memories: - memory_content = mem.get("memory_content", "") - message_time = mem.get("message_time", "") - when_to_use = mem.get("when_to_use", "") if self.enable_when_to_use else "" - metadata = {} - try: - metadata["time_int"] = int(message_time.split(" ")[0].replace("-", "")) - except Exception: - logger.warning(f"Invalid message time format: {message_time}") - - update_dicts.append( - { - "memory_id": mem.get("memory_id", ""), - "content": memory_content, - "when_to_use": when_to_use, - "message_time": message_time, - "author": self.author, - "metadata": metadata, - }, - ) - - if update_dicts: - handler = MemoryHandler(target, self.service_context) - memory_nodes = await handler.update_batch(update_dicts) - all_memory_nodes.extend(memory_nodes) - - if not all_memory_nodes: - return "No valid memories provided." - - self.memory_nodes.extend(all_memory_nodes) - output = f"Successfully updated {len(all_memory_nodes)} memories." - logger.info(output) - return output diff --git a/reme/memory/vector_tools/record/update_memory_v1.py b/reme/memory/vector_tools/record/update_memory_v1.py deleted file mode 100644 index 5d1551e7..00000000 --- a/reme/memory/vector_tools/record/update_memory_v1.py +++ /dev/null @@ -1,212 +0,0 @@ -"""Update memory in vector store""" - -from loguru import logger - -from .memory_handler import MemoryHandler -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall - - -class UpdateMemoryV1(BaseMemoryTool): - """Tool to update memories in vector store by deleting and adding memory entries""" - - def __init__( - self, - name="update_memory", - enable_memory_target: bool = False, - enable_when_to_use: bool = False, - **kwargs, - ): - kwargs["enable_multiple"] = True - super().__init__(name=name, **kwargs) - self.enable_memory_target: bool = enable_memory_target - self.enable_when_to_use: bool = enable_when_to_use - - def _build_memory_parameters(self, include_memory_id: bool = False) -> dict: - """Build the memory parameters schema based on enabled features. - - Args: - include_memory_id: If True, include memory_id field (for updates) - """ - properties = {} - required = [] - - if include_memory_id: - properties["memory_id"] = { - "type": "string", - "description": "ID of the memory to update", - } - required.append("memory_id") - - properties.update( - { - "message_time": { - "type": "string", - "description": "message time, e.g. '2020-01-01 00:00:00'", - }, - "memory_content": { - "type": "string", - "description": "content of the memory.", - }, - }, - ) - required.extend(["message_time", "memory_content"]) - - if self.enable_when_to_use: - properties["when_to_use"] = { - "type": "string", - "description": "description of when to use this memory.", - } - required.append("when_to_use") - - if self.enable_memory_target: - properties["memory_target"] = { - "type": "string", - "description": "target memory type for this memory.", - } - required.append("memory_target") - - return { - "type": "object", - "properties": properties, - "required": required, - } - - def _build_multiple_tool_call(self) -> ToolCall: - """Build and return the multiple tool call schema""" - return ToolCall( - **{ - "description": "update memories by updating existing memories and adding new memory entries.", - "parameters": { - "type": "object", - "properties": { - "memories_to_update": { - "type": "array", - "description": "List of memories to update", - "items": self._build_memory_parameters(include_memory_id=True), - }, - "memories_to_add": { - "type": "array", - "description": "List of memories to add", - "items": self._build_memory_parameters(include_memory_id=False), - }, - }, - "required": ["memories_to_update", "memories_to_add"], - }, - }, - ) - - async def execute(self): - # Get parameters - memories_to_update = self.context.get("memories_to_update", []) - memories_to_add = self.context.get("memories_to_add", []) - - if not memories_to_update and not memories_to_add: - return "No memories to update or add, operation completed." - - # Step 1: Collect and delete all old memories that need to be updated - if memories_to_update: - # Group deletion IDs by memory_target if enabled - if self.enable_memory_target: - delete_by_target = {} - for mem in memories_to_update: - target = mem.get("memory_target", self.memory_target) - memory_id = mem.get("memory_id") - if memory_id: - if target not in delete_by_target: - delete_by_target[target] = [] - delete_by_target[target].append(memory_id) - else: - delete_by_target = { - self.memory_target: [mem.get("memory_id") for mem in memories_to_update if mem.get("memory_id")], - } - - # Delete old memories for each target - for target, memory_ids in delete_by_target.items(): - if memory_ids: - handler = MemoryHandler(target, self.service_context) - await handler.delete(memory_ids) - - # Step 2: Prepare all memories to add (both updated and new) - all_memories_to_add = [] - - # Add memories from updates - if memories_to_update: - for mem in memories_to_update: - target = ( - mem.get( - "memory_target", - self.memory_target, - ) - if self.enable_memory_target - else self.memory_target - ) - all_memories_to_add.append((target, mem)) - - # Add new memories - if memories_to_add: - for mem in memories_to_add: - target = ( - mem.get( - "memory_target", - self.memory_target, - ) - if self.enable_memory_target - else self.memory_target - ) - all_memories_to_add.append((target, mem)) - - # Step 3: Group all memories by target and add them in batch - memories_by_target = {} - for target, mem in all_memories_to_add: - if target not in memories_by_target: - memories_by_target[target] = [] - memories_by_target[target].append(mem) - - # Process each target and add memories - all_memory_nodes = [] - updated_count = len(memories_to_update) - added_count = len(memories_to_add) - - for target, target_memories in memories_by_target.items(): - # Prepare memory data for batch add - memory_dicts = [] - for mem in target_memories: - memory_content = mem.get("memory_content", "") - message_time = mem.get("message_time", "") - when_to_use = mem.get("when_to_use", "") if self.enable_when_to_use else "" - metadata = {} - try: - metadata["time_int"] = int(message_time.split(" ")[0].replace("-", "")) - except Exception: - logger.warning(f"Invalid message time format: {message_time}") - - memory_dicts.append( - { - "content": memory_content, - "when_to_use": when_to_use, - "message_time": message_time, - "ref_memory_id": self.history_id, - "author": self.author, - "metadata": metadata, - }, - ) - - # Batch add all memories for this target - if memory_dicts: - handler = MemoryHandler(target, self.service_context) - memory_nodes = await handler.add_batch(memory_dicts) - all_memory_nodes.extend(memory_nodes) - - # Extend memory_nodes for tracking - self.memory_nodes.extend(all_memory_nodes) - - # Build output message - operations = [] - if updated_count > 0: - operations.append(f"updated {updated_count} memories.") - if added_count > 0: - operations.append(f"added {added_count} new memories.") - operations.append("Operation completed.") - logger.info("\n".join(operations)) - return "\n".join(operations) diff --git a/reme/memory/vector_tools/record/update_memory_v2.py b/reme/memory/vector_tools/record/update_memory_v2.py deleted file mode 100644 index 5b8f0950..00000000 --- a/reme/memory/vector_tools/record/update_memory_v2.py +++ /dev/null @@ -1,157 +0,0 @@ -"""Update memory in vector store""" - -from loguru import logger - -from .memory_handler import MemoryHandler -from ..base_memory_tool import BaseMemoryTool -from ....core.schema import ToolCall - - -class UpdateMemoryV2(BaseMemoryTool): - """Tool to update memories in vector store by deleting and adding memory entries""" - - def __init__( - self, - name="update_memory", - enable_memory_target: bool = False, - enable_when_to_use: bool = False, - **kwargs, - ): - kwargs["enable_multiple"] = True - super().__init__(name=name, **kwargs) - self.enable_memory_target: bool = enable_memory_target - self.enable_when_to_use: bool = enable_when_to_use - - def _build_add_memory_parameters(self) -> dict: - """Build the add memory parameters schema based on enabled features.""" - properties = { - "message_time": { - "type": "string", - "description": "message time, e.g. '2020-01-01 00:00:00'", - }, - "memory_content": { - "type": "string", - "description": "content of the memory.", - }, - } - required = ["message_time", "memory_content"] - - if self.enable_when_to_use: - properties["when_to_use"] = { - "type": "string", - "description": "description of when to use this memory.", - } - required.append("when_to_use") - - if self.enable_memory_target: - properties["memory_target"] = { - "type": "string", - "description": "target memory type for this memory.", - } - required.append("memory_target") - - return { - "type": "object", - "properties": properties, - "required": required, - } - - def _build_multiple_tool_call(self) -> ToolCall: - """Build and return the multiple tool call schema""" - return ToolCall( - **{ - "description": "update memories by removing and adding memory entries.", - "parameters": { - "type": "object", - "properties": { - "memory_ids_to_delete": { - "type": "array", - "description": "List of memory IDs to delete", - "items": { - "type": "string", - }, - }, - "memories_to_add": { - "type": "array", - "description": "List of memories to add", - "items": self._build_add_memory_parameters(), - }, - }, - "required": ["memory_ids_to_delete", "memories_to_add"], - }, - }, - ) - - async def execute(self): - # Get parameters - memory_ids_to_delete = self.context.get("memory_ids_to_delete", []) - memory_ids_to_delete = sorted({mid for mid in memory_ids_to_delete if mid}) - memories_to_add = self.context.get("memories_to_add", []) - - if not memory_ids_to_delete and not memories_to_add: - return "No memories to remove or add, operation completed." - - # Group memories by memory_target if enabled - if self.enable_memory_target: - memories_by_target = {} - for mem in memories_to_add: - target = mem.get("memory_target", self.memory_target) - if target not in memories_by_target: - memories_by_target[target] = [] - memories_by_target[target].append(mem) - else: - memories_by_target = {self.memory_target: memories_to_add} - - # Delete memories (all at once, regardless of target) - removed_count = 0 - if memory_ids_to_delete: - # Use the default memory_target handler for deletion - handler = MemoryHandler(self.memory_target, self.service_context) - await handler.delete(memory_ids_to_delete) - removed_count = len(memory_ids_to_delete) - - # Add new memories by target - added_count = 0 - all_memory_nodes = [] - for target, target_memories in memories_by_target.items(): - # Parse and prepare add data - add_dicts = [] - for mem in target_memories: - memory_content = mem.get("memory_content", "") - message_time = mem.get("message_time", "") - when_to_use = mem.get("when_to_use", "") if self.enable_when_to_use else "" - metadata = {} - try: - metadata["time_int"] = int(message_time.split(" ")[0].replace("-", "")) - except Exception: - logger.warning(f"Invalid message time format: {message_time}") - - add_dicts.append( - { - "content": memory_content, - "when_to_use": when_to_use, - "message_time": message_time, - "ref_memory_id": self.history_id, - "author": self.author, - "metadata": metadata, - }, - ) - - if add_dicts: - handler = MemoryHandler(target, self.service_context) - memory_nodes = await handler.add_batch(add_dicts) - all_memory_nodes.extend(memory_nodes) - added_count += len(memory_nodes) - - # Extend memory_nodes for tracking - self.memory_nodes.extend(all_memory_nodes) - - # Build output message - operations = [] - if removed_count > 0: - operations.append(f"removed {removed_count} old memories.") - if added_count > 0: - operations.append(f"added {added_count} new memories.") - operations.append("Operation completed.") - logger.info("\n".join(operations)) - return "\n".join(operations) diff --git a/reme/reme.py b/reme/reme.py index c7c04d2e..1a27d9f3 100644 --- a/reme/reme.py +++ b/reme/reme.py @@ -1,823 +1,47 @@ -"""ReMe classes for simplified configuration and execution.""" +"""ReMe memory management application entry point.""" +import asyncio import sys -from pathlib import Path -from .config import ReMeConfigParser -from .core import Application -from .core.enumeration import MemoryType, Role -from .core.schema import Message, MemoryNode -from .memory.vector_tools import ( - AddDraftAndRetrieveSimilarMemory, - AddHistory, - AddMemory, - DelegateTask, - ReadAllProfiles, - ReadHistory, - RetrieveProfile, - RetrieveMemory, - UpdateProfilesV1, -) -from .memory.vector_tools.profiles.profile_handler import ProfileHandler -from .memory.vector_tools.record.memory_handler import MemoryHandler -from .memory.vector_based import ( - BaseMemoryAgent, - PersonalRetriever, - PersonalSummarizer, - ProceduralRetriever, - ProceduralSummarizer, - ReMeRetriever, - ReMeSummarizer, - ToolRetriever, - ToolSummarizer, -) +from .application import Application +from .components import R +from .config import parse_args, resolve_app_config +from .enumeration import ComponentEnum +from .utils import cli_find_reme, load_env, precheck_start + +_CLIENT_KWARGS = {"host", "port", "timeout", "transport", "command", "args"} class ReMe(Application): - """ReMe with config file support and flow execution methods.""" - - def __init__( - self, - *args, - llm_api_key: str | None = None, - llm_base_url: str | None = None, - embedding_api_key: str | None = None, - embedding_base_url: str | None = None, - working_dir: str = ".reme", - config_path: str = "vector", - enable_logo: bool = True, - log_to_console: bool = True, - log_to_file: bool = True, - default_llm_config: dict | None = None, - default_embedding_model_config: dict | None = None, - default_vector_store_config: dict | None = None, - default_token_counter_config: dict | None = None, - target_user_names: list[str] | None = None, - target_task_names: list[str] | None = None, - target_tool_names: list[str] | None = None, - enable_profile: bool = True, - profile_backend: str = "filesystem", - profile_store_name: str = "profile", - profile_collection_name: str | None = None, - profile_max_capacity: int = 50, - **kwargs, - ): - """Initialize ReMe with config. - - Example: - ```python - reme = ReMe(...) - await reme.start() - await reme.summarize_memory(...) - await reme.retrieve_memory(...) - await reme.close() - ``` - - Args: - *args: Positional arguments forwarded to the base `Application`. - llm_api_key: API key used by the default LLM backend when provided. - llm_base_url: Base URL used by the default LLM backend when provided. - embedding_api_key: API key used by the default embedding backend when provided. - embedding_base_url: Base URL used by the default embedding backend when provided. - working_dir: Directory for generated config, logs, caches, and local stores. - config_path: Built-in config name or config file path used to initialize services. - enable_logo: Whether to print the ReMe logo during startup. - log_to_console: Whether to emit logs to the console. - log_to_file: Whether to write logs under `working_dir`. - default_llm_config: Overrides for the default LLM configuration. - default_embedding_model_config: Overrides for the default embedding model configuration. - default_vector_store_config: Configuration for the default memory vector store. - Its `collection_name` is used for normal memory storage. - default_token_counter_config: Overrides for the default token counter configuration. - target_user_names: Personal memory targets to register at initialization. - target_task_names: Procedural memory targets to register at initialization. - target_tool_names: Tool memory targets to register at initialization. - enable_profile: Whether to enable profile functionality. Set to False when using - profile-free memory flows. - profile_backend: Profile storage backend. Use "filesystem" for local JSONL profile - files or "vector" for a dedicated profile vector collection. - profile_store_name: Internal vector store key used to register and look up the - profile vector store in `service_context.vector_stores`. This is not the - database collection name. - profile_collection_name: Dedicated database collection/table name for vector - profiles. When unset, vector profiles use the default memory collection name - with a "_profile" suffix. - profile_max_capacity: Maximum number of profile rows to keep per memory target. - When the limit is exceeded, the oldest profile rows are removed. - **kwargs: Additional keyword arguments forwarded to the base `Application`. - """ - super().__init__( - *args, - llm_api_key=llm_api_key, - llm_base_url=llm_base_url, - embedding_api_key=embedding_api_key, - embedding_base_url=embedding_base_url, - working_dir=working_dir, - config_path=config_path, - enable_logo=enable_logo, - log_to_console=log_to_console, - log_to_file=log_to_file, - parser=ReMeConfigParser, - default_llm_config=default_llm_config, - default_embedding_model_config=default_embedding_model_config, - default_vector_store_config=default_vector_store_config, - default_token_counter_config=default_token_counter_config, - **kwargs, - ) - - self.enable_profile = enable_profile - self.profile_backend = profile_backend - self.profile_store_name = profile_store_name - self.profile_collection_name = profile_collection_name - self.profile_max_capacity = profile_max_capacity - - memory_target_type_mapping: dict[str, MemoryType] = {} - if target_user_names: - for name in target_user_names: - assert name not in memory_target_type_mapping, f"target_user_names={name} is already used." - memory_target_type_mapping[name] = MemoryType.PERSONAL - - if target_task_names: - for name in target_task_names: - assert name not in memory_target_type_mapping, f"target_task_names={name} is already used." - memory_target_type_mapping[name] = MemoryType.PROCEDURAL - - if target_tool_names: - for name in target_tool_names: - assert name not in memory_target_type_mapping, f"target_tool_names={name} is already used." - memory_target_type_mapping[name] = MemoryType.TOOL - - self.service_context.memory_target_type_mapping = memory_target_type_mapping - - if self.enable_profile and self.profile_backend == "filesystem": - profile_path = Path(self.service_context.service_config.working_dir) / "profile" - profile_path.mkdir(parents=True, exist_ok=True) - self.profile_dir: str = str(profile_path) - else: - self.profile_dir: str = "" - - if self.enable_profile and self.profile_backend == "vector": - self._ensure_profile_vector_store_config() - - def _add_meta_memory(self, memory_type: str | MemoryType, memory_target: str): - """Register or validate a memory target with the given memory type.""" - if memory_target in self.service_context.memory_target_type_mapping: - assert self.service_context.memory_target_type_mapping[memory_target] is memory_type - else: - self.service_context.memory_target_type_mapping[memory_target] = MemoryType(memory_type) - - @staticmethod - def _resolve_memory_target( - user_name: str = "", - task_name: str = "", - tool_name: str = "", - ) -> tuple[MemoryType, str]: - """Resolve memory type and target from user_name, task_name, or tool_name. - - Args: - user_name: User name for personal memory - task_name: Task name for procedural memory - tool_name: Tool name for tool memory - - Returns: - tuple: (memory_type, memory_target) - - Raises: - RuntimeError: If none or multiple memory targets are specified - """ - if user_name: - memory_type = MemoryType.PERSONAL - memory_target = user_name - assert not task_name and not tool_name, "Cannot add task and tool memory when user memory is specified" - - elif task_name: - memory_type = MemoryType.PROCEDURAL - memory_target = task_name - assert not user_name and not tool_name, "Cannot add user and tool memory when task memory is specified" - - elif tool_name: - memory_type = MemoryType.TOOL - memory_target = tool_name - assert not user_name and not task_name, "Cannot add user and task memory when tool memory is specified" - - else: - raise RuntimeError("Must specify user_name, task_name, or tool_name") - - return memory_type, memory_target - - def _ensure_started(self) -> None: - """Ensure memory operations run only after services are initialized.""" - if not self._started: - raise RuntimeError("ReMe is not started. Call `await reme.start()` before using memory APIs.") - - @staticmethod - def _unwrap_memory_result( - result: str | dict, - operation_name: str, - return_dict: bool, - ) -> str | dict: - """Normalize memory API results and fail loudly on swallowed inner errors.""" - if not isinstance(result, dict): - raise RuntimeError(f"{operation_name} failed before producing a structured result: {result}") - - if "answer" not in result: - raise RuntimeError(f"{operation_name} returned an invalid result payload: missing 'answer'") - - if return_dict: - return result - return result["answer"] - - def _ensure_profile_vector_store_config(self) -> None: - """Ensure the dedicated profile vector store exists in service config.""" - vector_store_configs = self.service_context.service_config.vector_stores - if "default" not in vector_store_configs: - raise RuntimeError("Vector profile backend requires a default vector store configuration") - - default_config = vector_store_configs["default"] - profile_collection_name = self.profile_collection_name or f"{default_config.collection_name}_profile" - - if self.profile_store_name in vector_store_configs: - if self.profile_collection_name: - vector_store_configs[self.profile_store_name] = vector_store_configs[ - self.profile_store_name - ].model_copy( - update={"collection_name": profile_collection_name}, - ) - return - - vector_store_configs[self.profile_store_name] = default_config.model_copy( - update={"collection_name": profile_collection_name}, - ) - - def _get_profile_tool_kwargs(self, raise_exception: bool) -> dict: - """Shared profile tool configuration.""" - return { - "profile_dir": self.profile_dir, - "profile_backend": self.profile_backend, - "profile_store_name": self.profile_store_name, - "profile_max_capacity": self.profile_max_capacity, - "raise_exception": raise_exception, - } - - async def summarize_memory( - self, - messages: list[Message | dict], - description: str = "", - user_name: str | list[str] = "", - task_name: str | list[str] = "", - tool_name: str | list[str] = "", - enable_thinking_params: bool = True, - version: str = "default", - retrieve_top_k: int = 20, - return_dict: bool = False, - raise_exception: bool = False, - llm_config_name: str = "default", - **kwargs, - ) -> str | dict: - """Summarize personal, procedural and tool memories for the given context.""" - self._ensure_started() - format_messages: list[Message] = [] - for message in messages: - if isinstance(message, dict): - assert message.get("time_created"), "message must have time_created field." - message = Message(**message) - format_messages.append(message) - - if version == "default": - profile_tool_kwargs = self._get_profile_tool_kwargs(raise_exception) - personal_summarizer_tools: list = [ - AddDraftAndRetrieveSimilarMemory( - enable_thinking_params=enable_thinking_params, - enable_memory_target=False, - enable_when_to_use=False, - enable_multiple=True, - top_k=retrieve_top_k, - raise_exception=raise_exception, - ), - AddMemory( - enable_thinking_params=enable_thinking_params, - enable_memory_target=False, - enable_when_to_use=False, - enable_multiple=True, - raise_exception=raise_exception, - ), - ] - if self.enable_profile: - if self.profile_backend == "vector": - profile_context_tool = RetrieveProfile( - top_k=min(5, retrieve_top_k), - enable_thinking_params=False, - enable_memory_target=False, - enable_multiple=False, - **profile_tool_kwargs, - ) - else: - profile_context_tool = ReadAllProfiles( - enable_thinking_params=False, - enable_memory_target=False, - **profile_tool_kwargs, - ) - personal_summarizer_tools.extend( - [ - profile_context_tool, - UpdateProfilesV1( - enable_thinking_params=enable_thinking_params, - enable_memory_target=False, - enable_multiple=True, - **profile_tool_kwargs, - ), - ], - ) - personal_summarizer: BaseMemoryAgent = PersonalSummarizer( - llm=llm_config_name, - tools=personal_summarizer_tools, - raise_exception=raise_exception, - ) - - else: - raise NotImplementedError(f"version={version} is not supported") - - procedural_summarizer: BaseMemoryAgent = ProceduralSummarizer( - llm=llm_config_name, - tools=[ - AddDraftAndRetrieveSimilarMemory( - enable_thinking_params=enable_thinking_params, - enable_memory_target=False, - enable_when_to_use=False, - enable_multiple=True, - top_k=retrieve_top_k, - raise_exception=raise_exception, - ), - AddMemory( - enable_thinking_params=enable_thinking_params, - enable_memory_target=False, - enable_when_to_use=False, - enable_multiple=True, - raise_exception=raise_exception, - ), - ], - raise_exception=raise_exception, - ) - tool_summarizer: BaseMemoryAgent = ToolSummarizer( - llm=llm_config_name, - tools=[ - AddDraftAndRetrieveSimilarMemory( - enable_thinking_params=enable_thinking_params, - enable_memory_target=False, - enable_when_to_use=False, - enable_multiple=True, - top_k=retrieve_top_k, - raise_exception=raise_exception, - ), - AddMemory( - enable_thinking_params=enable_thinking_params, - enable_memory_target=False, - enable_when_to_use=False, - enable_multiple=True, - raise_exception=raise_exception, - ), - ], - raise_exception=raise_exception, - ) - - memory_agents = [] - memory_targets = [] - if user_name: - if isinstance(user_name, str): - for message in format_messages: - if message.role is Role.USER: - message.name = user_name - self._add_meta_memory(MemoryType.PERSONAL, user_name) - memory_targets.append(user_name) - elif isinstance(user_name, list): - for name in user_name: - self._add_meta_memory(MemoryType.PERSONAL, name) - memory_targets.append(name) - else: - raise RuntimeError("user_name must be str or list[str]") - memory_agents.append(personal_summarizer) - - if task_name: - if isinstance(task_name, str): - self._add_meta_memory(MemoryType.PROCEDURAL, task_name) - memory_targets.append(task_name) - elif isinstance(task_name, list): - for name in task_name: - self._add_meta_memory(MemoryType.PROCEDURAL, name) - memory_targets.append(name) - else: - raise RuntimeError("task_name must be str or list[str]") - memory_agents.append(procedural_summarizer) - - if tool_name: - if isinstance(tool_name, str): - self._add_meta_memory(MemoryType.TOOL, tool_name) - memory_targets.append(tool_name) - elif isinstance(tool_name, list): - for name in tool_name: - self._add_meta_memory(MemoryType.TOOL, name) - memory_targets.append(name) - else: - raise RuntimeError("tool_name must be str or list[str]") - memory_agents.append(tool_summarizer) - - if not memory_agents: - memory_agents = [personal_summarizer, procedural_summarizer, tool_summarizer] - - reme_summarizer: BaseMemoryAgent = ReMeSummarizer( - tools=[ - AddHistory(raise_exception=raise_exception), - DelegateTask(memory_agents=memory_agents, raise_exception=raise_exception), - ], - raise_exception=raise_exception, - ) - - result = await reme_summarizer.call( - messages=format_messages, - description=description, - service_context=self.service_context, - memory_targets=memory_targets, - **kwargs, - ) - - return self._unwrap_memory_result(result, "summarize_memory", return_dict) - - async def retrieve_memory( - self, - query: str = "", - description: str = "", - messages: list[dict] | None = None, - user_name: str | list[str] = "", - task_name: str | list[str] = "", - tool_name: str | list[str] = "", - enable_thinking_params: bool = True, - version: str = "default", - retrieve_top_k: int = 20, - enable_time_filter: bool = True, - return_dict: bool = False, - raise_exception: bool = False, - llm_config_name: str = "default", - **kwargs, - ) -> str | dict: - """Retrieve relevant personal, procedural and tool memories for a query.""" - self._ensure_started() - - if version == "default": - profile_tool_kwargs = self._get_profile_tool_kwargs(raise_exception) - personal_retriever_tools = [] - if self.enable_profile: - if self.profile_backend == "vector": - profile_context_tool = RetrieveProfile( - top_k=min(5, retrieve_top_k), - enable_thinking_params=False, - enable_memory_target=False, - enable_multiple=False, - **profile_tool_kwargs, - ) - else: - profile_context_tool = ReadAllProfiles( - enable_thinking_params=False, - enable_memory_target=False, - **profile_tool_kwargs, - ) - personal_retriever_tools.append(profile_context_tool) - personal_retriever_tools.extend( - [ - RetrieveMemory( - top_k=retrieve_top_k, - enable_thinking_params=enable_thinking_params, - enable_time_filter=enable_time_filter, - enable_multiple=True, - raise_exception=raise_exception, - ), - ReadHistory( - enable_thinking_params=enable_thinking_params, - enable_multiple=True, - raise_exception=raise_exception, - ), - ], - ) - personal_retriever: BaseMemoryAgent = PersonalRetriever( - llm=llm_config_name, - tools=personal_retriever_tools, - raise_exception=raise_exception, - ) - else: - raise NotImplementedError(f"version={version} is not supported") - - procedural_retriever: BaseMemoryAgent = ProceduralRetriever( - llm=llm_config_name, - tools=[ - RetrieveMemory( - top_k=retrieve_top_k, - enable_thinking_params=enable_thinking_params, - enable_time_filter=False, - enable_multiple=True, - raise_exception=raise_exception, - ), - ReadHistory( - enable_thinking_params=enable_thinking_params, - enable_multiple=True, - raise_exception=raise_exception, - ), - ], - raise_exception=raise_exception, - ) - tool_retriever: BaseMemoryAgent = ToolRetriever( - llm=llm_config_name, - tools=[ - RetrieveMemory( - top_k=retrieve_top_k, - enable_thinking_params=enable_thinking_params, - enable_time_filter=False, - enable_multiple=True, - raise_exception=raise_exception, - ), - ReadHistory( - enable_thinking_params=enable_thinking_params, - enable_multiple=True, - raise_exception=raise_exception, - ), - ], - raise_exception=raise_exception, - ) - - memory_agents = [] - memory_targets = [] - if user_name: - if isinstance(user_name, str): - self._add_meta_memory(MemoryType.PERSONAL, user_name) - memory_targets.append(user_name) - elif isinstance(user_name, list): - for name in user_name: - self._add_meta_memory(MemoryType.PERSONAL, name) - memory_targets.append(name) - else: - raise RuntimeError("user_name must be str or list[str]") - memory_agents.append(personal_retriever) - - if task_name: - if isinstance(task_name, str): - self._add_meta_memory(MemoryType.PROCEDURAL, task_name) - memory_targets.append(task_name) - elif isinstance(task_name, list): - for name in task_name: - self._add_meta_memory(MemoryType.PROCEDURAL, name) - memory_targets.append(name) - else: - raise RuntimeError("task_name must be str or list[str]") - memory_agents.append(procedural_retriever) - - if tool_name: - if isinstance(tool_name, str): - self._add_meta_memory(MemoryType.TOOL, tool_name) - memory_targets.append(tool_name) - elif isinstance(tool_name, list): - for name in tool_name: - self._add_meta_memory(MemoryType.TOOL, name) - memory_targets.append(name) - else: - raise RuntimeError("tool_name must be str or list[str]") - memory_agents.append(tool_retriever) - - if not memory_agents: - memory_agents = [personal_retriever, procedural_retriever, tool_retriever] - - reme_retriever: BaseMemoryAgent = ReMeRetriever( - tools=[DelegateTask(memory_agents=memory_agents, raise_exception=raise_exception)], - raise_exception=raise_exception, - ) - - result = await reme_retriever.call( - query=query, - messages=messages, - description=description, - service_context=self.service_context, - memory_targets=memory_targets, - **kwargs, - ) - - return self._unwrap_memory_result(result, "retrieve_memory", return_dict) - - async def retrieve_profile( - self, - query: str | list[str], - user_name: str, - top_k: int = 5, - return_dict: bool = False, - ) -> str | dict: - """Retrieve relevant profile rows for a user.""" - self._ensure_started() - if not self.enable_profile: - raise RuntimeError("Profile functionality is disabled.") - - profile_handler = self.get_profile_handler(user_name) - if profile_handler is None: - raise RuntimeError("Profile functionality is disabled.") - - retrieved_nodes, output = await profile_handler.aretrieve( - query=query, - limit=top_k, - add_profile_id=True, - add_history_id=True, - ) - result = { - "answer": output or "No matching profiles found.", - "retrieved_nodes": retrieved_nodes, - } - return self._unwrap_memory_result(result, "retrieve_profile", return_dict) - - async def add_memory( - self, - memory_content: str, - user_name: str = "", - task_name: str = "", - tool_name: str = "", - when_to_use: str = "", - message_time: str = "", - ref_memory_id: str = "", - author: str = "", - score: float = 0.0, - **kwargs, - ): - """Add memory to the vector store. - - Args: - memory_content: The content of the memory to add - user_name: User name for personal memory - task_name: Task name for procedural memory - tool_name: Tool name for tool memory - when_to_use: Description of when this memory should be used - message_time: Timestamp of the message - ref_memory_id: Reference to another memory ID - author: Author of the memory - score: Score/importance of the memory - **kwargs: Additional metadata - - Returns: - MemoryNode: The created memory node - """ - memory_type, memory_target = self._resolve_memory_target(user_name, task_name, tool_name) - self._add_meta_memory(memory_type, memory_target) - - handler = self.get_memory_handler(memory_target) - memory_node = await handler.add( - content=memory_content, - when_to_use=when_to_use, - message_time=message_time, - ref_memory_id=ref_memory_id, - author=author, - score=score, - **kwargs, - ) - return memory_node - - async def get_memory( - self, - memory_id: str, - ): - """Get a memory node by its memory_id. - - Args: - memory_id: The ID of the memory to retrieve - - Returns: - MemoryNode: The retrieved memory node - """ - vector_node = await self.default_vector_store.get(memory_id) - return MemoryNode.from_vector_node(vector_node) - - async def delete_memory( - self, - memory_id: str, - ): - """Delete a memory node by its memory_id. - - Args: - memory_id: The ID of the memory to delete - """ - await self.default_vector_store.delete(memory_id) - - async def delete_all(self): - """Delete all memory nodes in the vector store.""" - await self.default_vector_store.delete_all() - - async def update_memory( - self, - memory_id: str, - user_name: str = "", - task_name: str = "", - tool_name: str = "", - memory_content: str | None = None, - when_to_use: str | None = None, - message_time: str | None = None, - ref_memory_id: str | None = None, - author: str | None = None, - score: float | None = None, - **kwargs, - ): - """Update a memory node's content and/or metadata. - - Args: - memory_id: The ID of the memory to update - user_name: User name for personal memory - task_name: Task name for procedural memory - tool_name: Tool name for tool memory - memory_content: New content for the memory (optional) - when_to_use: New description of when to use (optional) - message_time: New timestamp (optional) - ref_memory_id: New reference memory ID (optional) - author: New author (optional) - score: New score/importance (optional) - **kwargs: Additional metadata to update - - Returns: - MemoryNode: The updated memory node - """ - memory_type, memory_target = self._resolve_memory_target(user_name, task_name, tool_name) - self._add_meta_memory(memory_type, memory_target) - - handler = self.get_memory_handler(memory_target) - memory_node = await handler.update( - memory_id=memory_id, - content=memory_content, - when_to_use=when_to_use, - message_time=message_time, - ref_memory_id=ref_memory_id, - author=author, - score=score, - **kwargs, - ) - return memory_node - - async def list_memory( - self, - user_name: str = "", - task_name: str = "", - tool_name: str = "", - filters: dict | None = None, - limit: int | None = None, - sort_key: str | None = None, - reverse: bool = True, - ): - """List memory nodes with optional filtering and sorting. - - Args: - user_name: User name for personal memory - task_name: Task name for procedural memory - tool_name: Tool name for tool memory - filters: Additional filters to apply (optional) - limit: Maximum number of results to return (optional) - sort_key: Field to sort by (optional) - reverse: Sort in reverse order (default: True) - - Returns: - list[MemoryNode]: List of memory nodes - """ - memory_type, memory_target = self._resolve_memory_target(user_name, task_name, tool_name) - self._add_meta_memory(memory_type, memory_target) - - handler = self.get_memory_handler(memory_target) - memory_nodes = await handler.list( - filters=filters, - limit=limit, - sort_key=sort_key, - reverse=reverse, - ) - return memory_nodes - - def get_memory_handler(self, memory_target: str) -> MemoryHandler: - """Get the memory handler for the specified memory target.""" - return MemoryHandler(memory_target=memory_target, service_context=self.service_context) - - @property - def profile_path(self) -> Path | None: - """Get the path to the profile directory. Returns None if profile is disabled.""" - if not self.enable_profile or self.profile_backend != "filesystem": - return None - collection_name = self.service_context.service_config.vector_stores["default"].collection_name - return Path(self.profile_dir) / collection_name - - def get_profile_handler(self, user_name: str) -> ProfileHandler | None: - """Get the profile handler for the specified user. Returns None if profile is disabled.""" - if not self.enable_profile: - return None - return ProfileHandler( - memory_target=user_name, - profile_path=self.profile_path, - service_context=self.service_context, - profile_backend=self.profile_backend, - profile_store_name=self.profile_store_name, - max_capacity=self.profile_max_capacity, - ) + """ReMe memory management application.""" + + +async def call_server(action: str, **kwargs): + """Call the appropriate server component.""" + backend: str = kwargs.pop("backend", "http") + client_kwargs = {key: kwargs.pop(key) for key in list(kwargs) if key in _CLIENT_KWARGS} + client_cls = R.get(ComponentEnum.CLIENT, backend) + if client_cls is None: + raise ValueError(f"Unknown client backend: {backend!r}") + async with client_cls(**client_kwargs) as client: + async for chunk in client(action=action, **kwargs): + print(chunk, end="", flush=True) + print() def main(): - """Main entry point for running ReMe from command line.""" - from . import extension # noqa: F401 # pylint: disable=unused-import - from . import memory # noqa: F401 # pylint: disable=unused-import - - ReMe(*sys.argv[1:], config_path="service").run_service() + """Parse CLI arguments and launch the appropriate mode.""" + action, kwargs = parse_args(*sys.argv[1:]) + if action == "start": + load_env() + kwargs = resolve_app_config(**kwargs) + if not precheck_start(kwargs.get("service")): + return + ReMe(**kwargs).run_app() + elif action == "find_reme": + cli_find_reme() + else: + asyncio.run(call_server(action, **kwargs)) if __name__ == "__main__": diff --git a/reme/reme_cli.py b/reme/reme_cli.py deleted file mode 100644 index 5d086532..00000000 --- a/reme/reme_cli.py +++ /dev/null @@ -1,192 +0,0 @@ -"""ReMe File System""" - -import asyncio -import sys -from pathlib import Path - -from prompt_toolkit import PromptSession - -from .config import ReMeConfigParser -from .core import Application - -from .core.utils import play_horse_easter_egg -from .memory.file_based.components import CliAgent - - -class ReMeCli(Application): - """ReMe Cli""" - - def __init__( - self, - *args, - working_dir: str = ".reme", - config_path: str = "cli", - enable_logo: bool = True, - log_to_console: bool = True, - llm_api_key: str | None = None, - llm_base_url: str | None = None, - embedding_api_key: str | None = None, - embedding_base_url: str | None = None, - default_as_llm_config: dict | None = None, - default_embedding_model_config: dict | None = None, - default_file_store_config: dict | None = None, - default_token_counter_config: dict | None = None, - default_file_watcher_config: dict | None = None, - context_window_tokens: int = 128000, - reserve_tokens: int = 36000, - keep_recent_tokens: int = 20000, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - **kwargs, - ): - """Initialize ReMe with config.""" - working_path = Path(working_dir) - working_path.mkdir(parents=True, exist_ok=True) - memory_path = working_path / "memory" - memory_path.mkdir(parents=True, exist_ok=True) - self.working_dir: str = str(working_path.absolute()) - - default_file_watcher_config = default_file_watcher_config or {} - if not default_file_watcher_config.get("watch_paths", None): - default_file_watcher_config["watch_paths"] = [ - str(working_path / "MEMORY.md"), - str(working_path / "memory.md"), - str(memory_path), - ] - super().__init__( - *args, - llm_api_key=llm_api_key, - llm_base_url=llm_base_url, - embedding_api_key=embedding_api_key, - embedding_base_url=embedding_base_url, - working_dir=working_dir, - config_path=config_path, - enable_logo=enable_logo, - log_to_console=log_to_console, - parser=ReMeConfigParser, - default_as_llm_config=default_as_llm_config, - default_embedding_model_config=default_embedding_model_config, - default_file_store_config=default_file_store_config, - default_token_counter_config=default_token_counter_config, - default_file_watcher_config=default_file_watcher_config, - **kwargs, - ) - - self.service_config.metadata.setdefault("context_window_tokens", context_window_tokens) - self.service_config.metadata.setdefault("reserve_tokens", reserve_tokens) - self.service_config.metadata.setdefault("keep_recent_tokens", keep_recent_tokens) - self.service_config.metadata.setdefault("vector_weight", vector_weight) - self.service_config.metadata.setdefault("candidate_multiplier", candidate_multiplier) - - self.commands = { - "/new": "Create a new conversation.", - "/compact": "Compact messages into a summary.", - "/exit": "Exit the application.", - "/clear": "Clear the history.", - "/help": "Show help.", - "/horse": "A surprise.", - } - - async def chat_with_remy(self, **kwargs): - """Interactive CLI chat with Remy using simple streaming output.""" - language = self.service_config.language - print(f"ReMe language={language or 'default'}") - - cli_agent = CliAgent( - vector_weight=self.service_config.metadata["vector_weight"], - candidate_multiplier=self.service_config.metadata["candidate_multiplier"], - context_window_tokens=self.service_config.metadata["context_window_tokens"], - reserve_tokens=self.service_config.metadata["reserve_tokens"], - keep_recent_tokens=self.service_config.metadata["keep_recent_tokens"], - working_dir=self.working_dir, - language=language, - **kwargs, - ) - session = PromptSession() - - # Print welcome banner - print("\n========================================") - print(" Welcome to Remy Chat!") - print("========================================\n") - - while True: - try: - # Get user input (async) - user_input = await session.prompt_async("You: ") - user_input = user_input.strip() - if not user_input: - continue - - # Handle commands - if user_input == "/exit": - break - - if user_input == "/new": - result = await cli_agent.new() - print(f"{result}\nConversation reset\n") - continue - - if user_input == "/compact": - result = await cli_agent.compact(force_compact=True) - print(f"{result}\nHistory compacted.\n") - continue - - if user_input == "/history": - result = cli_agent.format_history() - print(f"Formated History:\n{result}\n") - continue - - if user_input == "/clear": - cli_agent.messages.clear() - print("History cleared.\n") - continue - - if user_input == "/help": - print("\nCommands:") - for command, description in self.commands.items(): - print(f" {command}: {description}") - continue - - if user_input == "/horse": - play_horse_easter_egg() - continue - - try: - await cli_agent.call( - query=user_input, - service_context=self.service_context, - ) - except Exception as e: - print(f"\nStream error: {e}") - - # End current streaming line - print("\n") - print("----------------------------------------\n") - - except EOFError: - break - except KeyboardInterrupt: - print("\nInterrupted.") - break - except Exception as e: - print(f"Error: {e}") - import traceback - - traceback.print_exc() - - print("\nGoodbye!\n") - - -async def async_main(): - """Main function for testing the ReMeFs CLI.""" - async with ReMeCli(*sys.argv[1:], log_to_console=False) as reme: - await reme.chat_with_remy() - - -def main(): - """Main function for testing the ReMeFs CLI.""" - asyncio.run(async_main()) - - -if __name__ == "__main__": - main() diff --git a/reme/reme_light.py b/reme/reme_light.py deleted file mode 100644 index 60f0a6c9..00000000 --- a/reme/reme_light.py +++ /dev/null @@ -1,835 +0,0 @@ -""" -ReMe Light Application Module - -This module provides the ReMeLight class, a specialized application built on top of -ReMe's core Application framework. It integrates memory management capabilities -including memory compaction, summarization, tool result management, and semantic -memory search functionality. - -Key Features: - - Memory compaction and summarization for long conversations - - Tool result compaction with file-based storage for large outputs - - Semantic memory search using vector and full-text search - - Configurable embedding models and vector store backends - - Async task management for background summarization -""" - -import asyncio -from pathlib import Path - -from agentscope.formatter import FormatterBase -from agentscope.message import Msg, TextBlock -from agentscope.model import ChatModelBase -from agentscope.token import HuggingFaceTokenCounter -from agentscope.tool import Toolkit, ToolResponse - -from .config import ReMeConfigParser -from .core import Application -from .core.utils import get_logger -from .memory.file_based import ReMeInMemoryMemory -from .memory.file_based.components import ( - Compactor, - ContextChecker, - Summarizer, - ToolResultCompactor, -) -from .memory.file_based.tools import FileIO, MemorySearch -from .memory.file_based.utils import AsMsgHandler - -logger = get_logger() - - -class ReMeLight(Application): - """ - ReMe Light Application Class. - - A lightweight memory-enabled application that provides semantic search, - memory compaction, summarization, and tool result management capabilities. - Built on top of the core Application framework with integrated vector store - and file-based memory management. - - This class is designed for applications requiring: - - Long conversation memory management with automatic compaction - - Semantic search over stored memories using hybrid vector/text search - - Background summarization of conversation history - - Automatic cleanup of expired tool results - - Attributes: - working_path (Path): Absolute path to the working directory. - memory_path (Path): Path to the memory storage directory. - tool_result_path (Path): Path to the tool result storage directory. - dialog_path (Path): Path to the dialog storage directory for raw conversation records. - vector_weight (float): Weight for vector search in hybrid search (0-1). - candidate_multiplier (float): Multiplier for candidate retrieval count. - summary_tasks (list[asyncio.Task]): List of active background summary tasks. - """ - - def __init__( - self, - working_dir: str = ".reme", - llm_api_key: str | None = None, - llm_base_url: str | None = None, - embedding_api_key: str | None = None, - embedding_base_url: str | None = None, - default_as_llm_config: dict | None = None, - default_embedding_model_config: dict | None = None, - default_file_store_config: dict | None = None, - default_file_watcher_config: dict | None = None, - vector_weight: float = 0.7, - candidate_multiplier: float = 3.0, - enable_load_env: bool = False, - ): - """ - Initialize the ReMeLight application. - - Sets up the working directory structure, configures API connections, - and initializes memory management components. - - Args: - working_dir (str): Base directory for all application data storage. - Defaults to ".reme". Will be created if it doesn't exist. - llm_api_key (str | None): API key for the language model service. - If None, will attempt to use environment variables. - llm_base_url (str | None): Base URL for the language model API endpoint. - If None, will use the default endpoint. - embedding_api_key (str | None): API key for the embedding model service. - If None, will attempt to use environment variables. - embedding_base_url (str | None): Base URL for the embedding API endpoint. - If None, will use the default endpoint. - default_as_llm_config (dict | None): Default configuration dictionary - for AgentScope language model. Overrides default settings. - default_embedding_model_config (dict | None): Default configuration - dictionary for the embedding model. - default_file_store_config (dict | None): Default configuration - dictionary for the file storage backend. - default_file_watcher_config (dict | None): Default configuration - dictionary for the file watcher. If ``watch_paths`` is included, - it is used as-is. Otherwise the built-in watch paths (MEMORY.md, - memory.md, and the memory directory) are used. - vector_weight (float): Weight assigned to vector similarity search - in hybrid search operations. Range [0.0, 1.0], default 0.7. - Higher values prioritize semantic similarity over keyword matching. - candidate_multiplier (float): Multiplier applied to max_results when - retrieving candidates for re-ranking. Default 3.0 means 3x more - candidates are retrieved than the final result count. - enable_load_env (bool): Whether to load environment variables from - .env file. Defaults to False. - - Note: - The following directory structure will be created: - - {working_dir}/ - Root working directory - - {working_dir}/memory/ - Memory storage files - - {working_dir}/tool_results/ - Compacted tool result files - - {working_dir}/dialog/ - Raw conversation records - """ - # Initialize working directory structure - self.working_path = Path(working_dir).absolute() - self.working_path.mkdir(parents=True, exist_ok=True) - self.memory_path = self.working_path / "memory" - self.memory_path.mkdir(parents=True, exist_ok=True) - self.tool_result_path = self.working_path / "tool_results" - self.tool_result_path.mkdir(parents=True, exist_ok=True) - self.dialog_path = self.working_path / "dialog" - self.dialog_path.mkdir(parents=True, exist_ok=True) - - self.vector_weight: float = vector_weight - self.candidate_multiplier: float = candidate_multiplier - - # Pick the existing memory markdown file. On case-insensitive filesystems - # (Windows NTFS, macOS APFS/HFS+) ``MEMORY.md`` and ``memory.md`` are the - # same file, so this also avoids watching it twice. Default to ``MEMORY.md`` - # when neither exists yet. - _memory_md = self.working_path / "MEMORY.md" - if not _memory_md.exists() and (self.working_path / "memory.md").exists(): - _memory_md = self.working_path / "memory.md" - _default_watch_paths = [str(_memory_md), str(self.memory_path)] - if default_file_watcher_config and default_file_watcher_config.get("watch_paths"): - _merged_file_watcher_config = default_file_watcher_config - else: - _merged_file_watcher_config = { - **(default_file_watcher_config or {}), - "watch_paths": _default_watch_paths, - } - - # Initialize the parent Application class with comprehensive configuration - super().__init__( - llm_api_key=llm_api_key, - llm_base_url=llm_base_url, - embedding_api_key=embedding_api_key, - embedding_base_url=embedding_base_url, - working_dir=str(self.working_path), - config_path="light", - enable_logo=False, - log_to_console=False, - enable_load_env=enable_load_env, - parser=ReMeConfigParser, - default_as_llm_config=default_as_llm_config, - default_embedding_model_config=default_embedding_model_config, - default_file_store_config=default_file_store_config, - default_file_watcher_config=_merged_file_watcher_config, - ) - - # Initialize list to track background summarization tasks - self.summary_tasks: list[asyncio.Task] = [] - - @staticmethod - def calculate_memory_compact_threshold(max_input_length: float, compact_ratio: float) -> int: - """Calculate the memory compaction threshold based on input length and ratio. - - Args: - max_input_length: Maximum input length in tokens. - compact_ratio: Ratio of the input length to use as the threshold. - - Returns: - Computed compaction threshold as an integer. - """ - return int(max_input_length * compact_ratio * 0.95) - - def _cleanup_tool_results(self) -> int: - """ - Clean up expired tool result files from the tool result directory. - - This method removes tool result files that have exceeded the retention - period. It helps manage disk space by automatically removing old, unused - tool outputs. - - Returns: - int: The number of files that were successfully deleted - """ - try: - # Create a compactor instance with default configuration - compactor = ToolResultCompactor(tool_result_dir=self.tool_result_path) - # Execute cleanup and return count of deleted files - return compactor.cleanup_expired_files() - except Exception as e: - # Log exception details but return 0 to indicate failure gracefully - logger.exception(f"Error cleaning up tool results: {e}") - return 0 - - async def start(self): - """ - Start the application lifecycle. - - Initializes all application components by calling the parent class start - method, then performs initial cleanup of expired tool result files. - - Returns: - The result from the parent Application.start() method. - - Note: - This method should be called before using any other application - functionality. It ensures all services are properly initialized. - """ - result = await super().start() - # Perform initial cleanup of any expired tool result files - self._cleanup_tool_results() - return result - - async def close(self) -> bool: - """ - Close the application and perform cleanup. - - Performs final cleanup of expired tool result files and then shuts down - all application components by calling the parent class close method. - - Returns: - bool: True if the application was closed successfully, False otherwise. - - Note: - This method should be called when the application is no longer needed - to ensure proper resource cleanup and data persistence. - """ - # Final cleanup of expired tool result files before shutdown - self._cleanup_tool_results() - return await super().close() - - async def compact_tool_result( - self, - messages: list[Msg], - old_max_bytes: int = 3000, - recent_max_bytes: int = 100 * 1024, - retention_days: int = 3, - recent_n: int = 1, - ) -> list[Msg]: - """ - Compact tool results by truncating large outputs and saving full content to files. - - This method processes a list of messages containing tool results and compacts - any that exceed the configured threshold. Large tool outputs are truncated - in the message while the full content is saved to separate files for later - retrieval if needed. - - Args: - messages (list[Msg]): List of messages potentially containing tool results - that may need compaction. - old_max_bytes (int): Byte threshold for old (non-recent) messages. Default 3000. - recent_max_bytes (int): Byte threshold for recent messages (trailing consecutive - tool-result messages). Default 100KB (102400 bytes). Content exceeding this - limit is saved to disk; the message retains the first 100KB with a - read_file-style truncation notice and the saved file path. - retention_days (int): Number of days to retain tool result files. - Default 3. - recent_n (int): Minimum number of most-recent tool-result messages to treat - as "recent" (using recent_max_bytes). The actual recent window is the - larger of this value and the trailing consecutive tool-result run. - Default 1. - - Returns: - list[Msg]: The processed list of messages with large tool results compacted. - If an error occurs, returns the original unmodified messages. - - Note: - - Recent tool results (trailing consecutive tool-result messages) are truncated - to recent_max_bytes using read_file-style output with a file path hint. - - Old tool results are truncated to old_max_bytes bytes. - - Full content of truncated results is saved to tool_result_path. - - Expired files are automatically cleaned up during this operation. - """ - try: - # Create compactor with instance configuration - compactor = ToolResultCompactor( - tool_result_dir=self.tool_result_path, - retention_days=retention_days, - old_max_bytes=old_max_bytes, - recent_max_bytes=recent_max_bytes, - recent_n=recent_n, - ) - - # Execute compaction and get processed messages - result = await compactor.call(messages=messages, service_context=self.service_context) - - # Clean up any expired tool result files during compaction - compactor.cleanup_expired_files() - - return result - - except Exception as e: - # Log the error and return original messages to maintain functionality - logger.exception(f"Error compacting tool results: {e}") - return messages - - async def check_context( - self, - messages: list[Msg], - memory_compact_threshold: int, - memory_compact_reserve: int = 10000, - as_token_counter: str | HuggingFaceTokenCounter = "default", - ) -> tuple[list[Msg], list[Msg], bool]: - """ - Check context size and determine if compaction is needed. - - Analyzes the provided messages to determine if they exceed the configured - token threshold and splits them into two groups: messages that should be - compacted and messages to keep in context. - - Args: - messages (list[Msg]): List of messages to check for context overflow. - memory_compact_threshold (int): Token count threshold for triggering - compaction. Messages exceeding this threshold will be split. - memory_compact_reserve (int): Token count to reserve for recent messages - to keep in context. Defaults to 10000 tokens. - as_token_counter (str | HuggingFaceTokenCounter): The token counter to use. - - Returns: - tuple[list[Msg], list[Msg], bool]: A tuple containing: - - messages_to_compact (list[Msg]): Older messages that should - be compacted/summarized. - - messages_to_keep (list[Msg]): Recent messages to keep in context. - - is_valid (bool): True if the split is valid (tool calls aligned), - False if splitting would break conversation integrity. - - Note: - - Returns ([], messages, True) if no compaction is needed. - - Ensures conversation pairs (user-assistant) are not split. - - is_valid=False indicates tool_use and tool_result are misaligned. - """ - try: - checker = ContextChecker( - memory_compact_threshold=memory_compact_threshold, - memory_compact_reserve=memory_compact_reserve, - as_token_counter=as_token_counter, - ) - - return await checker.call( - messages=messages, - service_context=self.service_context, - ) - - except Exception as e: - logger.exception(f"Error checking context: {e}") - return [], messages, False - - async def compact_memory( - self, - messages: list[Msg], - as_llm: str | ChatModelBase = "default", - as_llm_formatter: str | FormatterBase = "default", - as_token_counter: str | HuggingFaceTokenCounter = "default", - language: str = "zh", - max_input_length: float = 128 * 1024, - compact_ratio: float = 0.7, - previous_summary: str = "", - return_dict: bool = False, - add_thinking_block: bool = True, - extra_instruction: str = "", - ) -> str | dict: - """ - Compact a list of messages into a condensed summary. - - Uses the configured language model to generate a concise summary of the - provided messages. This is useful for reducing context window usage while - preserving important information from the conversation history. - - Args: - messages (list[Msg]): List of messages to be compacted into a summary. - as_llm (str | ChatModelBase): Language model identifier or instance - to use for summarization. Defaults to "default". - as_llm_formatter (str | FormatterBase): Formatter for the language model. - Defaults to "default". - as_token_counter (str | HuggingFaceTokenCounter): Token counter for - measuring message length. Defaults to "default". - language (str): Language for the summary output. "zh" for Chinese, - any other value for English. Defaults to "zh". - max_input_length (float): Maximum input length in tokens for the model. - Defaults to 128K tokens. - compact_ratio (float): Ratio used to calculate compaction threshold. - Defaults to 0.7. - previous_summary (str): Previous summary to incorporate into the new - summary for continuity. Defaults to empty string. - return_dict (bool): If True, returns a dict with user_message, - history_compact, and is_valid. Defaults to False. - add_thinking_block (bool): If True, adds a thinking block to the summary. - extra_instruction (str): Optional additional instruction appended to the - compaction prompt. Use this to guide what information to keep or - remove. For example: "Remove debug logs and tool-call details. Keep - requirements, decisions, and pending tasks." Defaults to empty string - (no extra instruction, preserving default behavior). - - Returns: - str | dict: The condensed summary string, or a dict containing - user_message, history_compact, and is_valid if return_dict=True. - Returns empty string or dict with empty values if an error occurred. - """ - try: - compactor = Compactor( - memory_compact_threshold=self.calculate_memory_compact_threshold(max_input_length, compact_ratio), - as_llm=as_llm, - as_llm_formatter=as_llm_formatter, - as_token_counter=as_token_counter, - language=language if language == "zh" else "", - return_dict=return_dict, - add_thinking_block=add_thinking_block, - extra_instruction=extra_instruction, - ) - - return await compactor.call( - messages=messages, - previous_summary=previous_summary, - service_context=self.service_context, - ) - - except Exception as e: - # Log error and return appropriate empty result - logger.exception(f"Error compacting memory: {e}") - if return_dict: - return {"user_message": str(e), "history_compact": str(e), "is_valid": False} - return "" - - async def summary_memory( - self, - messages: list[Msg], - as_llm: str | ChatModelBase = "default", - as_llm_formatter: str | FormatterBase = "default", - as_token_counter: str | HuggingFaceTokenCounter = "default", - toolkit: Toolkit | None = None, - language: str = "zh", - max_input_length: float = 128 * 1024, - compact_ratio: float = 0.7, - timezone: str | None = None, - add_thinking_block: bool = True, - ) -> str: - """ - Generate a comprehensive summary of the given messages. - - Creates a detailed summary of the conversation history and persists it - to the memory directory as structured files. Unlike compact_memory, this - method produces more detailed summaries suitable for long-term storage. - - Args: - messages (list[Msg]): List of messages to summarize. - as_llm (str | ChatModelBase): Language model identifier or instance - for summarization. Defaults to "default". - as_llm_formatter (str | FormatterBase): Formatter for the language model. - Defaults to "default". - as_token_counter (str | HuggingFaceTokenCounter): Token counter for - measuring message length. Defaults to "default". - toolkit (Toolkit | None): Toolkit with file operations for persisting - summaries. If None, creates a default toolkit with read/write/edit. - language (str): Language for the summary output. "zh" for Chinese, - any other value for English. Defaults to "zh". - max_input_length (float): Maximum input length in tokens. - Defaults to 128K tokens. - compact_ratio (float): Ratio used to calculate compaction threshold. - Defaults to 0.7. - timezone (str | None): Timezone string for date formatting - (e.g., "America/Chicago"). Defaults to system local time if None. - - Returns: - str: The generated summary text, or an empty string if an error occurred. - - Note: - This method may write summary files to the memory_path directory - using the provided or default toolkit. - """ - try: - if toolkit is None: - toolkit = Toolkit() - file_io = FileIO(working_dir=str(self.working_path)) - toolkit.register_tool_function(file_io.read_file) - toolkit.register_tool_function(file_io.write_file) - toolkit.register_tool_function(file_io.edit_file) - - summarizer = Summarizer( - working_dir=str(self.working_path), - memory_dir=str(self.memory_path), - memory_compact_threshold=self.calculate_memory_compact_threshold(max_input_length, compact_ratio), - toolkit=toolkit, - as_llm=as_llm, - as_llm_formatter=as_llm_formatter, - as_token_counter=as_token_counter, - language=language if language == "zh" else "", - timezone=timezone, - add_thinking_block=add_thinking_block, - ) - - return await summarizer.call(messages=messages, service_context=self.service_context) - - except Exception as e: - logger.exception(f"Error summarizing memory: {e}") - return "" - - def add_async_summary_task(self, messages: list[Msg], **kwargs): - """ - Add an asynchronous summary task for the given messages. - - Creates a background task to generate a summary of the provided messages - without blocking the main execution flow. Completed tasks are automatically - cleaned up from the task list. - - Args: - messages (list[Msg]): List of messages to be summarized asynchronously. - **kwargs: Additional keyword arguments passed to summary_memory(). - Supported arguments include: - - as_llm: Language model identifier or instance - - as_llm_formatter: Formatter for the language model - - as_token_counter: Token counter instance - - toolkit: Toolkit for file operations - - language: Output language ("zh" or other) - - max_input_length: Maximum input token length - - compact_ratio: Compaction threshold ratio - - Note: - - Completed/failed/canceled tasks are cleaned up before adding new ones - - Task results and errors are logged automatically - - Use await_summary_tasks() to wait for all pending tasks to complete - """ - remaining_tasks = [] - for task in self.summary_tasks: - if task.done(): - if task.cancelled(): - logger.warning("Summary task was cancelled.") - continue - exc = task.exception() - if exc is not None: - logger.error(f"Summary task failed: {exc}") - else: - result = task.result() - logger.info(f"Summary task completed: {result}") - else: - remaining_tasks.append(task) - self.summary_tasks = remaining_tasks - - task = asyncio.create_task(self.summary_memory(messages=messages, **kwargs)) - self.summary_tasks.append(task) - - @property - def default_as_token_counter(self) -> HuggingFaceTokenCounter: - """ - Get the default token counter for the memory. - - Returns: - HuggingFaceTokenCounter: The default token counter instance. - """ - return self.service_context.as_token_counters["default"] - - async def pre_reasoning_hook( - self, - messages: list[Msg], - system_prompt: str = "", - compressed_summary: str = "", - as_llm: str | ChatModelBase = "default", - as_llm_formatter: str | FormatterBase = "default", - as_token_counter: str | HuggingFaceTokenCounter = "default", - toolkit: Toolkit | None = None, - language: str = "zh", - max_input_length: float = 128 * 1024, - compact_ratio: float = 0.7, - memory_compact_reserve: int = 10000, - enable_tool_result_compact: bool = True, - tool_result_compact_keep_n: int = 3, - ) -> tuple[list[Msg], str]: - """ - Hook called before reasoning to manage memory and context. - - This method is designed to be called before each reasoning step to ensure - the conversation context fits within model limits. It performs tool result - compaction, checks context size, and triggers memory compaction if needed. - - Args: - messages (list[Msg]): Current conversation messages to be processed. - system_prompt (str): System prompt that will be included in the context. - Used to calculate available space. Defaults to empty string. - compressed_summary (str): Existing compressed summary from previous - compactions. Defaults to empty string. - as_llm (str | ChatModelBase): Language model for compaction operations. - Defaults to "default". - as_llm_formatter (str | FormatterBase): Formatter for the language model. - Defaults to "default". - as_token_counter (str | HuggingFaceTokenCounter): Token counter for - measuring content length. Defaults to "default". - toolkit (Toolkit | None): Toolkit for file operations in summarization. - Defaults to None. - language (str): Language for generated summaries. Defaults to "zh". - max_input_length (float): Maximum context window size in tokens. - Defaults to 128K tokens. - compact_ratio (float): Ratio for calculating compaction threshold. - Defaults to 0.7. - memory_compact_reserve (int): Token count to reserve for new responses. - Defaults to 10000 tokens. - enable_tool_result_compact (bool): Whether to compact tool results. - Defaults to True. - tool_result_compact_keep_n (int): Number of recent messages to exclude - from tool result compaction. Defaults to 3. - - Returns: - tuple[list[Msg], str]: A tuple containing: - - list[Msg]: Messages to keep in context (maybe reduced) - - str: Updated compressed summary incorporating compacted messages - - Note: - - Automatically triggers background summarization for compacted messages - - Tool results in recent messages (keep_n) are not compacted - - Returns original messages unchanged if no compaction is needed - """ - msg_handler = AsMsgHandler(self.default_as_token_counter) - - system_token_count = await msg_handler.count_str_token(system_prompt) - compressed_token_count = await msg_handler.count_str_token(compressed_summary) - memory_compact_threshold = self.calculate_memory_compact_threshold(max_input_length, compact_ratio) - left_compact_threshold = memory_compact_threshold - (system_token_count + compressed_token_count) - logger.info(f"Left compact threshold: {left_compact_threshold}") - - if enable_tool_result_compact and tool_result_compact_keep_n > 0: - compact_msgs = messages[:-tool_result_compact_keep_n] - await self.compact_tool_result(compact_msgs) - - messages_to_compact, messages_to_keep, is_valid = await self.check_context( - messages=messages, - memory_compact_threshold=left_compact_threshold, - memory_compact_reserve=memory_compact_reserve, - as_token_counter=as_token_counter, - ) - - if not messages_to_compact: - return messages, compressed_summary - - if not is_valid: - logger.warning("Invalid messages to compact, skipping.") - return messages, compressed_summary - - self.add_async_summary_task( - messages=messages_to_compact, - as_llm=as_llm, - as_llm_formatter=as_llm_formatter, - as_token_counter=as_token_counter, - toolkit=toolkit, - language=language, - max_input_length=max_input_length, - compact_ratio=compact_ratio, - ) - - compressed_summary = await self.compact_memory( - messages=messages_to_compact, - as_llm=as_llm, - as_llm_formatter=as_llm_formatter, - as_token_counter=as_token_counter, - language=language, - max_input_length=max_input_length, - compact_ratio=compact_ratio, - previous_summary=compressed_summary, - ) - - return messages_to_keep, compressed_summary - - async def await_summary_tasks(self) -> str: - """ - Wait for all background summary tasks to complete and collect results. - - Blocks until all pending summary tasks in the task list have completed, - canceled, or failed. Collects status information from each task and - clears the task list after processing. - - Returns: - str: A concatenated string of status messages for all tasks, including: - - Completion confirmations with results - - Cancellation notices - - Error messages for failed tasks - - Note: - - This method will block if any tasks are still running - - All tasks are removed from summary_tasks after this call - - Task exceptions are logged but do not raise to the caller - - Use this before application shutdown to ensure all summaries complete - """ - result = "" - for task in self.summary_tasks: - if task.done(): - # Task has already completed, check its status - if task.cancelled(): - logger.warning("Summary task was cancelled.") - result += "Summary task was cancelled.\n" - else: - # Check if the task raised an exception - exc = task.exception() - if exc is not None: - logger.error(f"Summary task failed: {exc}") - result += f"Summary task failed: {exc}\n" - else: - # Task completed successfully, collect result - task_result = task.result() - logger.info(f"Summary task completed: {task_result}") - result += f"Summary task completed: {task_result}\n" - - else: - # Task is still running, wait for it to complete - try: - task_result = await task - logger.info(f"Summary task completed: {task_result}") - result += f"Summary task completed: {task_result}\n" - - except asyncio.CancelledError: - logger.warning("Summary task was cancelled while waiting.") - result += "Summary task was cancelled.\n" - - except Exception as e: - logger.exception(f"Summary task failed: {e}") - result += f"Summary task failed: {e}\n" - - # Clear the task list after processing all tasks - self.summary_tasks.clear() - return result - - async def memory_search(self, query: str, max_results: int = 5, min_score: float = 0.1) -> ToolResponse: - """ - Mandatory recall step: semantically search MEMORY.md + memory/*.md - (and optional session transcripts) before answering questions about - prior work, decisions, dates, people, preferences, or todos; returns - top snippets with path + lines. - - Args: - query (str): The semantic search query to find relevant memory snippets. - max_results (int): Maximum number of search results to return (optional), default 5. - min_score (float): Minimum similarity score threshold for results (optional), default 0.1. - - Returns: - ToolResponse: A ToolResponse containing the search results as text, - or an error message if the query is empty. - """ - # Validate query parameter - if not query: - return ToolResponse( - content=[ - TextBlock( - type="text", - text="Error: No query provided.", - ), - ], - ) - - # Validate and clamp max_results to valid range [1, 100] - if isinstance(max_results, int): - max_results = min(max(max_results, 1), 100) - - elif isinstance(max_results, str): - try: - max_results = min(max(int(max_results), 1), 100) - except ValueError: - max_results = 5 - else: - max_results = 5 - - # Validate and clamp min_score to valid range [0.001, 0.999] - if isinstance(min_score, (int, float)): - min_score = float(min(max(min_score, 0.001), 0.999)) - - elif isinstance(min_score, str): - try: - min_score = float(min(max(float(min_score), 0.001), 0.999)) - except ValueError: - min_score = 0.1 - - else: - min_score = 0.1 - - # Initialize memory search tool with configured weights - search_tool = MemorySearch( - vector_weight=self.vector_weight, - candidate_multiplier=self.candidate_multiplier, - ) - - # Execute the search with validated parameters - search_result = await search_tool.call( - query=query, - max_results=max_results, - min_score=min_score, - service_context=self.service_context, - ) - - # Return results wrapped in ToolResponse format - return ToolResponse( - content=[ - TextBlock( - type="text", - text=search_result, - ), - ], - ) - - def get_in_memory_memory(self, as_token_counter: HuggingFaceTokenCounter | None = None): - """ - Create and return an in-memory memory instance. - - Factory method to create a ReMeInMemoryMemory instance configured with - the specified token counter. This memory instance stores messages in RAM - during the session, and automatically persists them to dialog_path when - messages are compressed or cleared. - - Args: - as_token_counter (HuggingFaceTokenCounter): Token counter for - measuring content length in the memory. - - Returns: - ReMeInMemoryMemory: A new in-memory memory instance ready for use. - The instance is configured with self.dialog_path for persistence. - - Note: - - Messages are stored in RAM during active session - - When messages are compressed via mark_messages_compressed(), they - are persisted to {dialog_path}/{date}.jsonl files - - When clear_content() is called, all messages are persisted before - clearing from memory - """ - return ReMeInMemoryMemory( - token_counter=as_token_counter or self.default_as_token_counter, - dialog_path=str(self.dialog_path), - ) diff --git a/reme4/schema/__init__.py b/reme/schema/__init__.py similarity index 100% rename from reme4/schema/__init__.py rename to reme/schema/__init__.py diff --git a/reme4/schema/application_config.py b/reme/schema/application_config.py similarity index 100% rename from reme4/schema/application_config.py rename to reme/schema/application_config.py diff --git a/reme4/schema/emb_node.py b/reme/schema/emb_node.py similarity index 100% rename from reme4/schema/emb_node.py rename to reme/schema/emb_node.py diff --git a/reme4/schema/file_chunk.py b/reme/schema/file_chunk.py similarity index 100% rename from reme4/schema/file_chunk.py rename to reme/schema/file_chunk.py diff --git a/reme4/schema/file_front_matter.py b/reme/schema/file_front_matter.py similarity index 100% rename from reme4/schema/file_front_matter.py rename to reme/schema/file_front_matter.py diff --git a/reme4/schema/file_link.py b/reme/schema/file_link.py similarity index 100% rename from reme4/schema/file_link.py rename to reme/schema/file_link.py diff --git a/reme4/schema/file_node.py b/reme/schema/file_node.py similarity index 100% rename from reme4/schema/file_node.py rename to reme/schema/file_node.py diff --git a/reme4/schema/request.py b/reme/schema/request.py similarity index 100% rename from reme4/schema/request.py rename to reme/schema/request.py diff --git a/reme4/schema/response.py b/reme/schema/response.py similarity index 100% rename from reme4/schema/response.py rename to reme/schema/response.py diff --git a/reme4/schema/stream_chunk.py b/reme/schema/stream_chunk.py similarity index 100% rename from reme4/schema/stream_chunk.py rename to reme/schema/stream_chunk.py diff --git a/reme4/steps/__init__.py b/reme/steps/__init__.py similarity index 100% rename from reme4/steps/__init__.py rename to reme/steps/__init__.py diff --git a/reme4/steps/base_step.py b/reme/steps/base_step.py similarity index 100% rename from reme4/steps/base_step.py rename to reme/steps/base_step.py diff --git a/reme4/steps/channel/__init__.py b/reme/steps/channel/__init__.py similarity index 100% rename from reme4/steps/channel/__init__.py rename to reme/steps/channel/__init__.py diff --git a/reme4/steps/channel/channel_notify.py b/reme/steps/channel/channel_notify.py similarity index 100% rename from reme4/steps/channel/channel_notify.py rename to reme/steps/channel/channel_notify.py diff --git a/reme4/steps/channel/claim_channel.py b/reme/steps/channel/claim_channel.py similarity index 100% rename from reme4/steps/channel/claim_channel.py rename to reme/steps/channel/claim_channel.py diff --git a/reme4/steps/common/__init__.py b/reme/steps/common/__init__.py similarity index 100% rename from reme4/steps/common/__init__.py rename to reme/steps/common/__init__.py diff --git a/reme4/steps/common/add.py b/reme/steps/common/add.py similarity index 100% rename from reme4/steps/common/add.py rename to reme/steps/common/add.py diff --git a/reme4/steps/common/demo.py b/reme/steps/common/demo.py similarity index 100% rename from reme4/steps/common/demo.py rename to reme/steps/common/demo.py diff --git a/reme4/steps/common/health_check.py b/reme/steps/common/health_check.py similarity index 100% rename from reme4/steps/common/health_check.py rename to reme/steps/common/health_check.py diff --git a/reme4/steps/common/help.py b/reme/steps/common/help.py similarity index 100% rename from reme4/steps/common/help.py rename to reme/steps/common/help.py diff --git a/reme4/steps/common/llm_demo.py b/reme/steps/common/llm_demo.py similarity index 100% rename from reme4/steps/common/llm_demo.py rename to reme/steps/common/llm_demo.py diff --git a/reme4/steps/common/stream_demo.py b/reme/steps/common/stream_demo.py similarity index 100% rename from reme4/steps/common/stream_demo.py rename to reme/steps/common/stream_demo.py diff --git a/reme4/steps/common/stream_llm_demo.py b/reme/steps/common/stream_llm_demo.py similarity index 100% rename from reme4/steps/common/stream_llm_demo.py rename to reme/steps/common/stream_llm_demo.py diff --git a/reme4/steps/common/version.py b/reme/steps/common/version.py similarity index 100% rename from reme4/steps/common/version.py rename to reme/steps/common/version.py diff --git a/reme4/steps/evolve/__init__.py b/reme/steps/evolve/__init__.py similarity index 100% rename from reme4/steps/evolve/__init__.py rename to reme/steps/evolve/__init__.py diff --git a/reme4/steps/evolve/_evolve.py b/reme/steps/evolve/_evolve.py similarity index 100% rename from reme4/steps/evolve/_evolve.py rename to reme/steps/evolve/_evolve.py diff --git a/reme4/steps/evolve/auto_memory.py b/reme/steps/evolve/auto_memory.py similarity index 100% rename from reme4/steps/evolve/auto_memory.py rename to reme/steps/evolve/auto_memory.py diff --git a/reme4/steps/evolve/auto_resource.py b/reme/steps/evolve/auto_resource.py similarity index 100% rename from reme4/steps/evolve/auto_resource.py rename to reme/steps/evolve/auto_resource.py diff --git a/reme4/steps/evolve/dream/__init__.py b/reme/steps/evolve/dream/__init__.py similarity index 100% rename from reme4/steps/evolve/dream/__init__.py rename to reme/steps/evolve/dream/__init__.py diff --git a/reme4/steps/evolve/dream/extract.py b/reme/steps/evolve/dream/extract.py similarity index 100% rename from reme4/steps/evolve/dream/extract.py rename to reme/steps/evolve/dream/extract.py diff --git a/reme4/steps/evolve/dream/finish.py b/reme/steps/evolve/dream/finish.py similarity index 100% rename from reme4/steps/evolve/dream/finish.py rename to reme/steps/evolve/dream/finish.py diff --git a/reme4/steps/evolve/dream/integrate.py b/reme/steps/evolve/dream/integrate.py similarity index 100% rename from reme4/steps/evolve/dream/integrate.py rename to reme/steps/evolve/dream/integrate.py diff --git a/reme4/steps/evolve/dream/proactive.py b/reme/steps/evolve/dream/proactive.py similarity index 100% rename from reme4/steps/evolve/dream/proactive.py rename to reme/steps/evolve/dream/proactive.py diff --git a/reme4/steps/evolve/dream/schema.py b/reme/steps/evolve/dream/schema.py similarity index 100% rename from reme4/steps/evolve/dream/schema.py rename to reme/steps/evolve/dream/schema.py diff --git a/reme4/steps/evolve/dream/topics.py b/reme/steps/evolve/dream/topics.py similarity index 100% rename from reme4/steps/evolve/dream/topics.py rename to reme/steps/evolve/dream/topics.py diff --git a/reme4/steps/evolve/dream/utils.py b/reme/steps/evolve/dream/utils.py similarity index 100% rename from reme4/steps/evolve/dream/utils.py rename to reme/steps/evolve/dream/utils.py diff --git a/reme4/steps/file_io/__init__.py b/reme/steps/file_io/__init__.py similarity index 100% rename from reme4/steps/file_io/__init__.py rename to reme/steps/file_io/__init__.py diff --git a/reme4/steps/file_io/_daily_index.py b/reme/steps/file_io/_daily_index.py similarity index 100% rename from reme4/steps/file_io/_daily_index.py rename to reme/steps/file_io/_daily_index.py diff --git a/reme4/steps/file_io/_file_io.py b/reme/steps/file_io/_file_io.py similarity index 100% rename from reme4/steps/file_io/_file_io.py rename to reme/steps/file_io/_file_io.py diff --git a/reme4/steps/file_io/_path.py b/reme/steps/file_io/_path.py similarity index 100% rename from reme4/steps/file_io/_path.py rename to reme/steps/file_io/_path.py diff --git a/reme4/steps/file_io/daily_create.py b/reme/steps/file_io/daily_create.py similarity index 100% rename from reme4/steps/file_io/daily_create.py rename to reme/steps/file_io/daily_create.py diff --git a/reme4/steps/file_io/daily_list.py b/reme/steps/file_io/daily_list.py similarity index 100% rename from reme4/steps/file_io/daily_list.py rename to reme/steps/file_io/daily_list.py diff --git a/reme4/steps/file_io/daily_reindex.py b/reme/steps/file_io/daily_reindex.py similarity index 100% rename from reme4/steps/file_io/daily_reindex.py rename to reme/steps/file_io/daily_reindex.py diff --git a/reme4/steps/file_io/delete.py b/reme/steps/file_io/delete.py similarity index 100% rename from reme4/steps/file_io/delete.py rename to reme/steps/file_io/delete.py diff --git a/reme4/steps/file_io/edit.py b/reme/steps/file_io/edit.py similarity index 100% rename from reme4/steps/file_io/edit.py rename to reme/steps/file_io/edit.py diff --git a/reme4/steps/file_io/frontmatter_delete.py b/reme/steps/file_io/frontmatter_delete.py similarity index 100% rename from reme4/steps/file_io/frontmatter_delete.py rename to reme/steps/file_io/frontmatter_delete.py diff --git a/reme4/steps/file_io/frontmatter_read.py b/reme/steps/file_io/frontmatter_read.py similarity index 100% rename from reme4/steps/file_io/frontmatter_read.py rename to reme/steps/file_io/frontmatter_read.py diff --git a/reme4/steps/file_io/frontmatter_update.py b/reme/steps/file_io/frontmatter_update.py similarity index 100% rename from reme4/steps/file_io/frontmatter_update.py rename to reme/steps/file_io/frontmatter_update.py diff --git a/reme4/steps/file_io/list.py b/reme/steps/file_io/list.py similarity index 100% rename from reme4/steps/file_io/list.py rename to reme/steps/file_io/list.py diff --git a/reme4/steps/file_io/move.py b/reme/steps/file_io/move.py similarity index 100% rename from reme4/steps/file_io/move.py rename to reme/steps/file_io/move.py diff --git a/reme4/steps/file_io/read.py b/reme/steps/file_io/read.py similarity index 100% rename from reme4/steps/file_io/read.py rename to reme/steps/file_io/read.py diff --git a/reme4/steps/file_io/read_image.py b/reme/steps/file_io/read_image.py similarity index 100% rename from reme4/steps/file_io/read_image.py rename to reme/steps/file_io/read_image.py diff --git a/reme4/steps/file_io/stat.py b/reme/steps/file_io/stat.py similarity index 100% rename from reme4/steps/file_io/stat.py rename to reme/steps/file_io/stat.py diff --git a/reme4/steps/file_io/write.py b/reme/steps/file_io/write.py similarity index 100% rename from reme4/steps/file_io/write.py rename to reme/steps/file_io/write.py diff --git a/reme4/steps/index/__init__.py b/reme/steps/index/__init__.py similarity index 100% rename from reme4/steps/index/__init__.py rename to reme/steps/index/__init__.py diff --git a/reme4/steps/index/_change_batch.py b/reme/steps/index/_change_batch.py similarity index 100% rename from reme4/steps/index/_change_batch.py rename to reme/steps/index/_change_batch.py diff --git a/reme4/steps/index/_watch_rules.py b/reme/steps/index/_watch_rules.py similarity index 100% rename from reme4/steps/index/_watch_rules.py rename to reme/steps/index/_watch_rules.py diff --git a/reme4/steps/index/clear_store.py b/reme/steps/index/clear_store.py similarity index 100% rename from reme4/steps/index/clear_store.py rename to reme/steps/index/clear_store.py diff --git a/reme4/steps/index/init_changes.py b/reme/steps/index/init_changes.py similarity index 100% rename from reme4/steps/index/init_changes.py rename to reme/steps/index/init_changes.py diff --git a/reme4/steps/index/log_changes.py b/reme/steps/index/log_changes.py similarity index 100% rename from reme4/steps/index/log_changes.py rename to reme/steps/index/log_changes.py diff --git a/reme4/steps/index/node_search.py b/reme/steps/index/node_search.py similarity index 100% rename from reme4/steps/index/node_search.py rename to reme/steps/index/node_search.py diff --git a/reme4/steps/index/search.py b/reme/steps/index/search.py similarity index 100% rename from reme4/steps/index/search.py rename to reme/steps/index/search.py diff --git a/reme4/steps/index/traverse.py b/reme/steps/index/traverse.py similarity index 100% rename from reme4/steps/index/traverse.py rename to reme/steps/index/traverse.py diff --git a/reme4/steps/index/update_changes.py b/reme/steps/index/update_changes.py similarity index 100% rename from reme4/steps/index/update_changes.py rename to reme/steps/index/update_changes.py diff --git a/reme4/steps/index/watch_changes.py b/reme/steps/index/watch_changes.py similarity index 100% rename from reme4/steps/index/watch_changes.py rename to reme/steps/index/watch_changes.py diff --git a/reme4/steps/transfer/__init__.py b/reme/steps/transfer/__init__.py similarity index 100% rename from reme4/steps/transfer/__init__.py rename to reme/steps/transfer/__init__.py diff --git a/reme4/steps/transfer/download.py b/reme/steps/transfer/download.py similarity index 100% rename from reme4/steps/transfer/download.py rename to reme/steps/transfer/download.py diff --git a/reme4/steps/transfer/ingest.py b/reme/steps/transfer/ingest.py similarity index 100% rename from reme4/steps/transfer/ingest.py rename to reme/steps/transfer/ingest.py diff --git a/reme4/steps/transfer/upload.py b/reme/steps/transfer/upload.py similarity index 100% rename from reme4/steps/transfer/upload.py rename to reme/steps/transfer/upload.py diff --git a/reme4/utils/__init__.py b/reme/utils/__init__.py similarity index 100% rename from reme4/utils/__init__.py rename to reme/utils/__init__.py diff --git a/reme4/utils/agent_state_io.py b/reme/utils/agent_state_io.py similarity index 100% rename from reme4/utils/agent_state_io.py rename to reme/utils/agent_state_io.py diff --git a/reme4/utils/common_utils.py b/reme/utils/common_utils.py similarity index 99% rename from reme4/utils/common_utils.py rename to reme/utils/common_utils.py index f874180f..fd03eb65 100644 --- a/reme4/utils/common_utils.py +++ b/reme/utils/common_utils.py @@ -154,7 +154,7 @@ async def mock_reme_server( cmd: list[str] = [ sys.executable, "-m", - "reme4.reme", + "reme.reme", "start", f"service.host={host}", f"service.port={port}", diff --git a/reme4/utils/env_utils.py b/reme/utils/env_utils.py similarity index 100% rename from reme4/utils/env_utils.py rename to reme/utils/env_utils.py diff --git a/reme4/utils/jsonl_zst.py b/reme/utils/jsonl_zst.py similarity index 100% rename from reme4/utils/jsonl_zst.py rename to reme/utils/jsonl_zst.py diff --git a/reme4/utils/link_expansion.py b/reme/utils/link_expansion.py similarity index 100% rename from reme4/utils/link_expansion.py rename to reme/utils/link_expansion.py diff --git a/reme4/utils/logger_utils.py b/reme/utils/logger_utils.py similarity index 100% rename from reme4/utils/logger_utils.py rename to reme/utils/logger_utils.py diff --git a/reme4/utils/logo_utils.py b/reme/utils/logo_utils.py similarity index 100% rename from reme4/utils/logo_utils.py rename to reme/utils/logo_utils.py diff --git a/reme4/utils/service_utils.py b/reme/utils/service_utils.py similarity index 100% rename from reme4/utils/service_utils.py rename to reme/utils/service_utils.py diff --git a/reme4/utils/similarity_utils.py b/reme/utils/similarity_utils.py similarity index 100% rename from reme4/utils/similarity_utils.py rename to reme/utils/similarity_utils.py diff --git a/reme4/utils/token_utils.py b/reme/utils/token_utils.py similarity index 100% rename from reme4/utils/token_utils.py rename to reme/utils/token_utils.py diff --git a/reme4/utils/wikilink_handler.py b/reme/utils/wikilink_handler.py similarity index 100% rename from reme4/utils/wikilink_handler.py rename to reme/utils/wikilink_handler.py diff --git a/reme4/__init__.py b/reme4/__init__.py deleted file mode 100644 index e4db3de0..00000000 --- a/reme4/__init__.py +++ /dev/null @@ -1,26 +0,0 @@ -"""ReMe CLI package.""" - -__version__ = "0.4.0.0" - -from . import config -from . import constants -from . import enumeration -from . import schema -from . import steps -from . import utils -from .application import Application -from .components import BaseComponent -from .reme import ReMe - -__all__ = [ - "Application", - "BaseComponent", - "ReMe", - # submodules - "config", - "constants", - "enumeration", - "schema", - "steps", - "utils", -] diff --git a/reme4/config/__init__.py b/reme4/config/__init__.py deleted file mode 100644 index c2903189..00000000 --- a/reme4/config/__init__.py +++ /dev/null @@ -1,8 +0,0 @@ -"""Config""" - -from .config_parser import parse_args, resolve_app_config - -__all__ = [ - "parse_args", - "resolve_app_config", -] diff --git a/reme4/reme.py b/reme4/reme.py deleted file mode 100644 index 1a27d9f3..00000000 --- a/reme4/reme.py +++ /dev/null @@ -1,48 +0,0 @@ -"""ReMe memory management application entry point.""" - -import asyncio -import sys - -from .application import Application -from .components import R -from .config import parse_args, resolve_app_config -from .enumeration import ComponentEnum -from .utils import cli_find_reme, load_env, precheck_start - -_CLIENT_KWARGS = {"host", "port", "timeout", "transport", "command", "args"} - - -class ReMe(Application): - """ReMe memory management application.""" - - -async def call_server(action: str, **kwargs): - """Call the appropriate server component.""" - backend: str = kwargs.pop("backend", "http") - client_kwargs = {key: kwargs.pop(key) for key in list(kwargs) if key in _CLIENT_KWARGS} - client_cls = R.get(ComponentEnum.CLIENT, backend) - if client_cls is None: - raise ValueError(f"Unknown client backend: {backend!r}") - async with client_cls(**client_kwargs) as client: - async for chunk in client(action=action, **kwargs): - print(chunk, end="", flush=True) - print() - - -def main(): - """Parse CLI arguments and launch the appropriate mode.""" - action, kwargs = parse_args(*sys.argv[1:]) - if action == "start": - load_env() - kwargs = resolve_app_config(**kwargs) - if not precheck_start(kwargs.get("service")): - return - ReMe(**kwargs).run_app() - elif action == "find_reme": - cli_find_reme() - else: - asyncio.run(call_server(action, **kwargs)) - - -if __name__ == "__main__": - main() diff --git a/reme_ai/__init__.py b/reme_ai/__init__.py deleted file mode 100644 index a1353f0e..00000000 --- a/reme_ai/__init__.py +++ /dev/null @@ -1,33 +0,0 @@ -# pylint: disable=wrong-import-position -"""ReMe AI - A memory management framework for AI agents.""" - -import os - -os.environ["FLOW_APP_NAME"] = "ReMe" - -from . import agent # noqa: E402 -from . import config # noqa: E402 -from . import constants # noqa: E402 -from . import enumeration # noqa: E402 -from . import retrieve # noqa: E402 -from . import schema # noqa: E402 -from . import service # noqa: E402 -from . import summary # noqa: E402 -from . import utils # noqa: E402 -from . import vector_store # noqa: E402 -from .main import ReMeApp # noqa: E402 F401 - -__all__ = [ - "agent", - "config", - "constants", - "enumeration", - "retrieve", - "schema", - "service", - "summary", - "utils", - "vector_store", -] - -__version__ = "0.2.0.7" diff --git a/reme_ai/agent/__init__.py b/reme_ai/agent/__init__.py deleted file mode 100644 index 2094f344..00000000 --- a/reme_ai/agent/__init__.py +++ /dev/null @@ -1,14 +0,0 @@ -"""Agent module for ReAct and tool-based agent implementations. - -This module provides submodules for different types of agent operations: -- react: ReAct (Reasoning and Acting) agent implementations -- tools: Mock search tools for testing and demonstration -""" - -from . import react -from . import tools - -__all__ = [ - "react", - "tools", -] diff --git a/reme_ai/agent/react/__init__.py b/reme_ai/agent/react/__init__.py deleted file mode 100644 index 11adaed4..00000000 --- a/reme_ai/agent/react/__init__.py +++ /dev/null @@ -1,13 +0,0 @@ -"""ReAct agent operations module. - -This module provides ReAct (Reasoning and Acting) agent implementations for -answering user queries through iterative reasoning and search actions. -""" - -from .agentic_retrieve_op import AgenticRetrieveOp -from .simple_react_op import SimpleReactOp - -__all__ = [ - "AgenticRetrieveOp", - "SimpleReactOp", -] diff --git a/reme_ai/agent/react/agentic_retrieve_op.py b/reme_ai/agent/react/agentic_retrieve_op.py deleted file mode 100644 index aab28e1a..00000000 --- a/reme_ai/agent/react/agentic_retrieve_op.py +++ /dev/null @@ -1,294 +0,0 @@ -""" -Async React agent operator tailored for retrieval workflows. - -This module implements a ReAct (Reasoning + Acting) agent pattern that combines -language model reasoning with tool execution. The agent iteratively: -1. Processes and manages conversation context (compaction/compression) -2. Generates reasoning and tool calls via LLM -3. Executes tools (e.g., file search, reading) -4. Incorporates tool results back into the conversation -5. Repeats until a final answer is reached or max_steps is exceeded - -The agent is specifically designed for RAG (Retrieval-Augmented Generation) workflows, -providing context management capabilities to handle long conversations efficiently. - -Context management is controlled via ``working_summary_mode`` and -``compact_ratio_threshold`` parameters, which are forwarded to -``MessageOffloadOp``. ``working_summary_mode`` selects between: -- ``compact`` – only compact verbose tool messages by storing full content externally - and keeping short previews in the context. -- ``compress`` – only apply LLM-based compression to generate a compact state snapshot. -- ``auto`` – first run compaction, then optionally run compression if the - compaction ratio is not sufficient (default). - -``compact_ratio_threshold`` is only used in ``auto`` mode and defines the compaction -ratio (tokens after compaction divided by original tokens) above which an additional -LLM-based compression pass is applied. It defaults to ``0.75``. -""" - -from typing import Dict, List - -from flowllm.core.context import C -from flowllm.core.enumeration import Role -from flowllm.core.op import BaseAsyncToolOp -from flowllm.core.schema import ToolCall, Message -from loguru import logger - - -@C.register_op() -class AgenticRetrieveOp(BaseAsyncToolOp): - """ - ReAct agent that exposes RAG-friendly tools and context policies. - - This agent implements the ReAct pattern where the LLM alternates between: - - Reasoning: Analyzing the problem and deciding what to do - - Acting: Executing tools to gather information - - Observing: Processing tool results and updating understanding - - The agent supports sophisticated context management to handle long conversations - through compaction (storing large tool messages externally) and compression - (LLM-based summarization of message history). - - Context management behavior is configured via ``working_summary_mode`` and - ``compact_ratio_threshold`` (see module docstring for details). These options are - passed to ``MessageOffloadOp`` to control whether the agent only compacts tool - messages, only compresses history, or applies an automatic compaction-then- - compression pipeline. - - Available tools: - - GrepOp: Search for patterns in files - - ReadFileOp: Read file contents - - The agent automatically manages context size by offloading large messages and - compressing conversation history when token limits are approached. - """ - - file_path: str = __file__ - - def __init__( - self, - llm: str = "qwen3_30b_instruct", - max_steps: int = 20, - **kwargs, - ): - """ - Initialize the agent runtime configuration. - - Args: - llm: Identifier for the language model to use. Defaults to "qwen3_30b_instruct". - max_steps: Maximum number of reasoning-action cycles before stopping. - Each cycle includes: context management -> LLM reasoning -> tool execution. - Defaults to 5 steps. - **kwargs: Additional arguments passed to the base BaseAsyncToolOp class. - """ - super().__init__(llm=llm, **kwargs) - # Maximum number of ReAct iterations (reasoning + tool execution cycles) - self.max_steps: int = max_steps - - def build_tool_call(self) -> ToolCall: - """ - Expose metadata describing how to invoke the agent. - - This method defines the tool schema that other components use to invoke - this agent. It specifies all input parameters including conversation messages - and context management configuration. - - Returns: - ToolCall: A schema object describing the agent's interface, including - all parameters for message handling and context management. - """ - return ToolCall( - **{ - "description": "A React agent that answers user queries.", - "input_schema": { - "messages": { - "type": "array", - "description": "messages", - "required": True, - }, - "working_summary_mode": { - "type": "string", - "description": "summary strategy: 'compact' only compacts large tool messages, 'compress' " - "only applies LLM-based compression, 'auto' first compacts then optionally compresses when " - "reduction is insufficient. Defaults to 'auto'.", - "required": False, - "enum": ["compact", "compress", "auto"], - }, - "compact_ratio_threshold": { - "type": "number", - "description": "Only used in 'auto' mode. Threshold for compaction (tokens after compaction " - "divided by original tokens). When the ratio is greater than this value, an additional " - "LLM-based compression pass is triggered. Defaults to 0.75.", - "required": False, - }, - "max_total_tokens": { - "type": "integer", - "description": "Maximum token threshold for triggering compression/compaction. For compaction " - "this is total tokens; for compression this excludes keep_recent_count and " - "system messages. Defaults to 20000.", - "required": False, - }, - "max_tool_message_tokens": { - "type": "integer", - "description": "Maximum token count per tool message before compaction applies. Exceeding " - "messages store full content externally with a preview in context. Defaults " - "to 2000.", - "required": False, - }, - "group_token_threshold": { - "type": "integer", - "description": "Maximum tokens per compression group for LLM-based compression. None/0 " - "compresses all messages together. Oversized messages form their own group. " - "Used in 'compress' or 'auto' mode.", - "required": False, - }, - "keep_recent_count": { - "type": "integer", - "description": "Number of recent messages preserved without compression/compaction. Defaults " - "to 1 for compaction and 2 for compression.", - "required": False, - }, - "store_dir": { - "type": "string", - "description": "Directory for storing offloaded contents. Required for compaction/compression " - "to save full tool messages and compressed groups.", - "required": False, - }, - "chat_id": { - "type": "string", - "description": "Chat session identifier for naming stored files. Defaults to auto-generated " - "UUID if omitted.", - "required": False, - }, - }, - }, - ) - - async def async_execute(self): - """ - Main execution loop implementing the ReAct (Reasoning + Acting) pattern. - - This method orchestrates the iterative agent workflow: - 1. Context Management: Compact/compress conversation history if needed - 2. Reasoning: LLM analyzes context and decides on actions (may include tool calls) - 3. Acting: Execute requested tools in parallel - 4. Observing: Incorporate tool results back into conversation - 5. Repeat until final answer (no tool calls) or max_steps reached - - The loop continues until: - - The LLM produces a response without tool calls (final answer) - - Maximum number of steps (max_steps) is reached - - Each iteration is called a "round" and represents one reasoning-action cycle. - """ - # Import tool operators that the agent can use - from reme_ai.retrieve.working import GrepOp, ReadFileOp, BatchWriteFileOp - from reme_ai.summary.working import MessageOffloadOp - - # Initialize available tools for the agent - # GrepOp: Search for patterns/text in files (useful for code search) - grep_op = GrepOp(language=self.language) - # ReadFileOp: Read and return file contents - read_file_op = ReadFileOp(language=self.language) - - # Create a dictionary mapping tool names to their operator instances - # This allows quick lookup when the LLM requests a specific tool - tool_op_dict: Dict[str, BaseAsyncToolOp] = { - grep_op.tool_call.name: grep_op, - read_file_op.tool_call.name: read_file_op, - } - - # Convert input messages to Message objects for processing - messages = [Message(**x) for x in self.context.messages] - - # Extract context management parameters from input, excluding messages - # These will be passed to the context management pipeline - context_kwargs = self.input_dict.copy() - context_kwargs.pop("messages", None) - - # Main ReAct loop: iterate up to max_steps times - for i in range(self.max_steps): - # Step 1: Context Management Phase - # Create a pipeline: MessageOffloadOp (compacts/compresses) -> BatchWriteFileOp (saves offloaded content) - # The >> operator chains these operations together - op = MessageOffloadOp() >> BatchWriteFileOp() - - # Apply context management to current message history - # This may compact large tool messages or compress old messages based on context_manage_mode - await op.async_call(messages=[x.simple_dump() for x in messages], **context_kwargs) - - # Update messages with the processed/optimized version from context management - # Large messages may now reference external files instead of containing full content - logger.info(f"round{i + 1}.offload={op.context.response.answer}") - messages = [Message(**x) for x in op.context.response.answer] - - # Step 2: Reasoning Phase - # LLM analyzes the conversation context and decides what to do - # It may generate a final answer or request tool calls to gather more information - assistant_message: Message = await self.llm.achat( - messages=messages, # Current conversation history (possibly optimized) - tools=[op.tool_call for op in tool_op_dict.values()], # Available tools the LLM can use - ) - - # Add the LLM's response to the conversation history - messages.append(assistant_message) - logger.info(f"round{i + 1}.assistant={assistant_message.model_dump_json()}") - - # Step 3: Check if we have a final answer - # If the LLM didn't request any tools, it has provided a final answer - # Exit the loop as the agent's work is complete - if not assistant_message.tool_calls: - break - - # Step 4: Acting Phase - Execute requested tools - # The LLM requested one or more tools to gather information - # We'll execute them in parallel for efficiency - op_list: List[BaseAsyncToolOp] = [] - - # Process each tool call requested by the LLM - for j, tool_call in enumerate(assistant_message.tool_calls): - # Validate that the requested tool exists in our available tools - if tool_call.name not in tool_op_dict: - logger.exception(f"unknown tool_call.name={tool_call.name}") - continue - - logger.info(f"round{i + 1}.{j} submit tool_calls={tool_call.name} argument={tool_call.argument_dict}") - - # Create a copy of the tool operator for this specific tool call - # Each tool call needs its own operator instance with the correct call ID - op_copy: BaseAsyncToolOp = tool_op_dict[tool_call.name].copy() - op_copy.tool_call.id = tool_call.id # Match the ID from LLM's tool call request - op_list.append(op_copy) - - # Submit the tool execution as an async task (runs in parallel with other tools) - # The op_copy instance is used to execute the tool with the specific call ID and arguments - self.submit_async_task(op_copy.async_call, **tool_call.argument_dict) - - # Wait for all submitted tool executions to complete - # This ensures we have all results before proceeding - await self.join_async_task() - - # Step 5: Observing Phase - Incorporate tool results into conversation - # Process each completed tool execution and add results to message history - for j, op in enumerate(op_list): - # Extract the tool execution result as a string - tool_result = str(op.output) - - # Create a tool message with the result, linked to the original tool call via ID - # This allows the LLM to associate results with the specific tool calls it made - tool_message = Message(role=Role.TOOL, content=tool_result, tool_call_id=op.tool_call.id) - messages.append(tool_message) - - # Log the tool result (truncated to first 200 chars for readability) - logger.info(f"round{i + 1}.{j} join tool_result={tool_result[:200]}...\n\n") - - # Loop continues: next iteration will process these tool results, - # manage context again if needed, and let LLM reason about the new information - - # After loop completes, set the final output - # The last message should be the LLM's final answer (without tool calls) - self.set_output(messages[-1].content) - - # Store the complete conversation history in the context response - # This includes all reasoning steps, tool calls, and tool results - self.context.response.metadata["messages"] = [x.simple_dump(add_reasoning=True) for x in messages] diff --git a/reme_ai/agent/react/simple_react_op.py b/reme_ai/agent/react/simple_react_op.py deleted file mode 100644 index 400267e1..00000000 --- a/reme_ai/agent/react/simple_react_op.py +++ /dev/null @@ -1,70 +0,0 @@ -"""Simple ReAct operation module. - -This module provides a simple ReAct (Reasoning and Acting) agent implementation -that extends the base ReactSearchOp for answering user queries through iterative -reasoning and search actions. -""" - -import asyncio - -from flowllm.core.context import C, FlowContext -from flowllm.gallery.agent import ReactSearchOp - - -@C.register_op() -class SimpleReactOp(ReactSearchOp): - """A simple ReAct (Reasoning and Acting) agent operation. - - This operation extends ReactSearchOp to provide a straightforward implementation - of a ReAct agent that answers user queries by reasoning about the problem and - taking search actions iteratively until a final answer is reached. - - The agent inherits all functionality from ReactSearchOp, including: - - Iterative reasoning and action cycles - - Search tool integration - - Maximum step limits for preventing infinite loops - """ - - def __init__( - self, - llm: str = "default", - max_steps: int = 20, - tool_call_interval: float = 1.0, - add_think_tool: bool = False, - **kwargs, - ): - """Initialize the agent runtime configuration.""" - super().__init__( - llm=llm, - max_steps=max_steps, - tool_call_interval=tool_call_interval, - add_think_tool=add_think_tool, - **kwargs, - ) - - -async def main(): - """Main function to demonstrate SimpleReactOp usage. - - This function initializes the FlowLLM context with ReMe configuration, - creates a SimpleReactOp instance, and processes a sample query about - stock prices for Maotai and Wuliangye. - - Example: - Run this module directly to test the SimpleReactOp: - ```bash - python -m reme_ai.agent.react.simple_react_op - ``` - """ - from reme_ai.config.config_parser import ConfigParser - - C.set_service_config(parser=ConfigParser, config_name="config=default").init_by_service_config() - context = FlowContext(query="茅台和五粮现在股价多少?") - - op = SimpleReactOp() - await op.async_call(context=context) - print(context.response.answer) - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/reme_ai/agent/tools/__init__.py b/reme_ai/agent/tools/__init__.py deleted file mode 100644 index 014482d2..00000000 --- a/reme_ai/agent/tools/__init__.py +++ /dev/null @@ -1,17 +0,0 @@ -"""Mock search tools for testing and demonstration purposes. - -This module provides mock search operations that simulate different search tool behaviors, -including LLM-based query classification and result generation. -""" - -from .llm_mock_search_op import LLMMockSearchOp -from .mock_search_tools import SearchToolA, SearchToolB, SearchToolC -from .use_mock_search_op import UseMockSearchOp - -__all__ = [ - "LLMMockSearchOp", - "SearchToolA", - "SearchToolB", - "SearchToolC", - "UseMockSearchOp", -] diff --git a/reme_ai/agent/tools/llm_mock_search_op.py b/reme_ai/agent/tools/llm_mock_search_op.py deleted file mode 100644 index 91313e25..00000000 --- a/reme_ai/agent/tools/llm_mock_search_op.py +++ /dev/null @@ -1,328 +0,0 @@ -"""LLM-based mock search operation for simulating search tool behavior. - -This module provides a mock search operation that uses LLM to classify queries -and generate realistic search results based on query complexity levels. -""" - -import asyncio -import json -import random -from typing import Dict, Any - -from flowllm.core.context import C, FlowContext -from flowllm.core.enumeration.role import Role -from flowllm.core.op import BaseAsyncToolOp -from flowllm.core.schema import Message -from flowllm.core.schema import ToolCall -from loguru import logger - - -@C.register_op() -class LLMMockSearchOp(BaseAsyncToolOp): - """ - Mock search operation that uses LLM to classify queries and simulate different scenarios. - - Supports three query complexity levels: - - simple: Simple factual queries with short, direct answers - - medium: Medium complexity queries requiring balanced performance - - complex: Complex research queries requiring comprehensive, in-depth results - - Each scenario can be configured with: - - success_rate: Probability of successful response (vs "Service busy" error) - - extra_time: Extra sleep time in seconds to simulate latency - - relevance_ratio: Probability of returning relevant results (vs random query results) - """ - - file_path: str = __file__ - - def __init__( - self, - llm: str = "qwen3_30b_instruct", - simple_config: Dict[str, Any] = None, - medium_config: Dict[str, Any] = None, - complex_config: Dict[str, Any] = None, - seed: int = 0, - **kwargs, - ): - """ - Initialize the LLM Mock Search Op. - - Args: - llm: LLM model name to use for classification and content generation - simple_config: Configuration for simple queries - - success_rate: float (0-1), default 0.95 - - extra_time: float (seconds), default 0.5 - - relevance_ratio: float (0-1), default 0.98 - medium_config: Configuration for medium complexity queries - - success_rate: float (0-1), default 0.85 - - extra_time: float (seconds), default 1.0 - - relevance_ratio: float (0-1), default 0.90 - complex_config: Configuration for complex queries - - success_rate: float (0-1), default 0.70 - - extra_time: float (seconds), default 1.5 - - relevance_ratio: float (0-1), default 0.80 - seed: Random seed for deterministic behavior, default 0 - """ - super().__init__(llm=llm, **kwargs) - - # Set random seed for deterministic behavior - self.seed = seed - random.seed(self.seed) - - # Default configurations for each scenario - self.simple_config = { - "success_rate": 0.95, - "extra_time": 0.5, - "relevance_ratio": 0.98, - "content_length": "short", - } - if simple_config: - self.simple_config.update(simple_config) - - self.medium_config = { - "success_rate": 0.85, - "extra_time": 1.0, - "relevance_ratio": 0.90, - "content_length": "medium", - } - if medium_config: - self.medium_config.update(medium_config) - - self.complex_config = { - "success_rate": 0.70, - "extra_time": 1.5, - "relevance_ratio": 0.80, - "content_length": "long", - } - if complex_config: - self.complex_config.update(complex_config) - - def build_tool_call(self) -> ToolCall: - """Build the tool call schema for the search operation. - - Returns: - ToolCall object defining the search tool interface - """ - return ToolCall( - **{ - "description": "Use search keywords to retrieve relevant information from the internet.", - "input_schema": { - "query": { - "type": "string", - "description": "search keyword or query", - "required": True, - }, - }, - }, - ) - - async def classify_query(self, query: str) -> str: - """ - Classify the query into simple, medium, or complex using LLM. - - Args: - query: The search query to classify - - Returns: - Classification result: "simple", "medium", or "complex" - """ - classification_prompt = self.prompt_format( - prompt_name="classification_prompt", - query=query, - ) - - messages = [Message(role=Role.USER, content=classification_prompt)] - - response = await self.llm.achat(messages=messages) - classification = response.content.strip().lower() - - # Extract classification from response - if "simple" in classification: - return "simple" - elif "complex" in classification: - return "complex" - else: - return "medium" - - async def generate_search_result(self, query: str, complexity: str, config: Dict[str, Any]) -> str: - """ - Generate mock search results using LLM based on query complexity. - - Args: - query: The search query - complexity: Query complexity level - config: Configuration for this complexity level - - Returns: - Generated search result content - """ - content_length = config["content_length"] - - generation_prompt = self.prompt_format( - prompt_name="generation_prompt", - query=query, - complexity=complexity, - content_length=content_length, - ) - - messages = [Message(role=Role.USER, content=generation_prompt)] - - response = await self.llm.achat(messages=messages) - return response.content - - async def generate_random_result(self) -> str: - """ - Generate a random/irrelevant search result to simulate low relevance. - - Returns: - Random search result content - """ - random_topics = [ - "the history of ancient civilizations", - "modern technology trends", - "climate change impacts", - "space exploration achievements", - "culinary traditions around the world", - "evolution of music genres", - "breakthroughs in medical science", - "architectural wonders", - "wildlife conservation efforts", - "developments in artificial intelligence", - ] - - random_query = random.choice(random_topics) - generation_prompt = self.prompt_format( - prompt_name="generation_prompt", - query=random_query, - complexity="simple", - content_length="short", - ) - - messages = [Message(role=Role.USER, content=generation_prompt)] - response = await self.llm.achat(messages=messages) - - return f"[Low Relevance Result]\n{response.content}" - - async def async_execute(self): - """Execute the mock search operation. - - This method classifies the query, applies the appropriate configuration, - simulates delays, and generates search results based on success and relevance rates. - """ - query: str = self.input_dict["query"] - logger.info(f"LLMMockSearchOp processing query: {query}") - - # Step 1: Classify the query - complexity = await self.classify_query(query) - logger.info(f"Query classified as: {complexity}") - - # Step 2: Get configuration for this complexity - if complexity == "simple": - config = self.simple_config - elif complexity == "medium": - config = self.medium_config - else: # complex - config = self.complex_config - - logger.info(f"Using config: {config}") - - # Step 3: Simulate extra time delay - extra_time = config["extra_time"] - await asyncio.sleep(extra_time) - logger.info(f"Simulated extra delay: {extra_time:.2f}s") - - # Step 4: Check success rate - if random.random() > config["success_rate"]: - error_message = "Search service is currently busy. Please try again later." - logger.warning(f"Simulated failure: {error_message}") - result_dict = { - "success": False, - "content": error_message, - "query": query, - "complexity": complexity, - } - self.set_output(json.dumps(result_dict, ensure_ascii=False)) - return - - # Step 5: Check relevance ratio - if random.random() > config["relevance_ratio"]: - # Generate random/irrelevant result - # NOTE: success=True because technically the tool executed without errors, - # but the content is irrelevant (low quality), which should result in score=0.0 during evaluation - logger.info("Generating low relevance result (success=True but low quality)") - content = await self.generate_random_result() - result_dict = { - "success": True, # Technical execution succeeded - "content": content, - "query": query, - "complexity": complexity, - "is_relevant": False, # Mark as irrelevant for debugging - } - else: - # Generate relevant result - logger.info("Generating relevant result") - content = await self.generate_search_result(query, complexity, config) - result_dict = { - "success": True, - "content": content, - "query": query, - "complexity": complexity, - "is_relevant": True, # Mark as relevant for debugging - } - - self.set_output(json.dumps(result_dict, ensure_ascii=False)) - - -async def async_main(): - """Main function for testing the LLMMockSearchOp with various query types.""" - from reme_ai.main import ReMeApp - - async with ReMeApp(): - # Test with different query types - test_queries = [ - "What is the capital of France?", # Simple - "How does quantum computing work?", # Medium - "Analyze the impact of artificial intelligence on global economy, employment, and society", # Complex - ] - - # Custom configurations for testing - custom_simple = { - "success_rate": 1, - "extra_time": 0, - "relevance_ratio": 1, - } - - custom_medium = { - "success_rate": 1, - "extra_time": 0, - "relevance_ratio": 1, - } - - custom_complex = { - "success_rate": 1, - "extra_time": 0, - "relevance_ratio": 1, - } - - op = LLMMockSearchOp( - simple_config=custom_simple, - medium_config=custom_medium, - complex_config=custom_complex, - ) - - for query in test_queries: - print(f"\n{'=' * 80}") - print(f"Testing query: {query}") - print(f"{'=' * 80}") - - context = FlowContext(query=query) - await op.async_call(context=context) - result = json.loads(context.llm_mock_search_result) - print(f"Success: {result['success']}") - print(f"Query: {result['query']}") - print(f"Complexity: {result['complexity']}") - print(f"Content:\n{result['content']}") - - -if __name__ == "__main__": - asyncio.run(async_main()) diff --git a/reme_ai/agent/tools/mock_search_tools.py b/reme_ai/agent/tools/mock_search_tools.py deleted file mode 100644 index 944011af..00000000 --- a/reme_ai/agent/tools/mock_search_tools.py +++ /dev/null @@ -1,182 +0,0 @@ -"""Specialized mock search tools with different performance characteristics. - -This module provides three search tools (SearchToolA, SearchToolB, SearchToolC) -each optimized for different query complexity levels, allowing for realistic -testing of tool selection strategies. -""" - -from flowllm.core.context import C -from flowllm.core.schema import ToolCall - -from reme_ai.agent.tools.llm_mock_search_op import LLMMockSearchOp - - -@C.register_op() -class SearchToolA(LLMMockSearchOp): - """Fast search tool optimized for simple queries. - - This tool is configured for quick responses with high success rates - on simple queries, but performs poorly on medium and complex queries. - Best suited for simple factual queries. - """ - - def __init__(self, llm: str = "qwen3_30b_instruct", **kwargs): - """Initialize SearchToolA with fast, simple-query-optimized configuration. - - Args: - llm: LLM model name to use - **kwargs: Additional arguments passed to LLMMockSearchOp - """ - # Configure for fast but shallow performance - simple_config = { - "success_rate": 0.9, # High success rate for simple queries - "extra_time": 0, # Very fast (0.2-0.5s range) - "relevance_ratio": 0.9, # High relevance - "content_length": "short", # Concise answers - } - - medium_config = { - "success_rate": 0.2, # Lower success for medium queries - "extra_time": 0, # Still fast - "relevance_ratio": 0.2, # Moderate relevance - "content_length": "short", # Limited depth - } - - complex_config = { - "success_rate": 0.5, # Poor success rate for complex queries - "extra_time": 0, # Fast but insufficient - "relevance_ratio": 0.5, # Low relevance (often misses key aspects) - "content_length": "short", # Too shallow for complex topics - } - - super().__init__( - llm=llm, - simple_config=simple_config, - medium_config=medium_config, - complex_config=complex_config, - **kwargs, - ) - - def build_tool_call(self) -> ToolCall: - """Build the tool call schema with description indicating simple query optimization. - - Returns: - ToolCall object with description indicating best use for simple queries - """ - tool_call = super().build_tool_call() - tool_call.description += " Best suited for simple queries." - return tool_call - - -@C.register_op() -class SearchToolB(LLMMockSearchOp): - """Balanced search tool optimized for medium complexity queries. - - This tool provides balanced performance across query types, with - excellent results for medium complexity queries. Best suited for - queries requiring moderate depth and context. - """ - - def __init__(self, llm: str = "qwen3_30b_instruct", **kwargs): - """Initialize SearchToolB with balanced, medium-query-optimized configuration. - - Args: - llm: LLM model name to use - **kwargs: Additional arguments passed to LLMMockSearchOp - """ - # Configure for balanced performance - simple_config = { - "success_rate": 0.3, # Very high success rate - "extra_time": 0, # Moderate speed (1.0-1.5s range) - "relevance_ratio": 0.3, # High relevance - "content_length": "medium", # More detailed than needed for simple - } - - medium_config = { - "success_rate": 0.9, # Excellent success rate - "extra_time": 0, # Balanced speed - "relevance_ratio": 0.9, # High relevance - "content_length": "medium", # Perfect depth for medium queries - } - - complex_config = { - "success_rate": 0.5, # Good success rate - "extra_time": 0, # Still reasonable speed - "relevance_ratio": 0.5, # Decent relevance but not exhaustive - "content_length": "medium", # Covers main points but lacks depth - } - - super().__init__( - llm=llm, - simple_config=simple_config, - medium_config=medium_config, - complex_config=complex_config, - **kwargs, - ) - - def build_tool_call(self) -> ToolCall: - """Build the tool call schema with description indicating medium query optimization. - - Returns: - ToolCall object with description indicating best use for medium complexity queries - """ - tool_call = super().build_tool_call() - tool_call.description += " Best suited for medium complexity queries." - return tool_call - - -@C.register_op() -class SearchToolC(LLMMockSearchOp): - """Comprehensive search tool optimized for complex queries. - - This tool provides thorough, in-depth results with high success rates - on complex queries, but may be slower and overly detailed for simple queries. - Best suited for complex research queries requiring comprehensive analysis. - """ - - def __init__(self, llm: str = "qwen3_30b_instruct", **kwargs): - """Initialize SearchToolC with comprehensive, complex-query-optimized configuration. - - Args: - llm: LLM model name to use - **kwargs: Additional arguments passed to LLMMockSearchOp - """ - # Configure for comprehensive but costly performance - simple_config = { - "success_rate": 0.3, # Good but not optimal (over-processing) - "extra_time": 0, # Slow (3.0-4.0s range) - "relevance_ratio": 0.3, # High relevance but unnecessary depth - "content_length": "long", # Too detailed for simple queries - } - - medium_config = { - "success_rate": 0.4, # High success rate - "extra_time": 0, # Slow but thorough - "relevance_ratio": 0.4, # High relevance with extra context - "content_length": "long", # More depth than needed - } - - complex_config = { - "success_rate": 0.9, # Excellent success rate - "extra_time": 0, # Slow but comprehensive (3.5-5.0s range) - "relevance_ratio": 0.9, # Very high relevance - "content_length": "long", # Perfect depth for complex queries - } - - super().__init__( - llm=llm, - simple_config=simple_config, - medium_config=medium_config, - complex_config=complex_config, - **kwargs, - ) - - def build_tool_call(self) -> ToolCall: - """Build the tool call schema with description indicating complex query optimization. - - Returns: - ToolCall object with description indicating best use for complex queries - """ - tool_call = super().build_tool_call() - tool_call.description += " Best suited for complex queries." - return tool_call diff --git a/reme_ai/agent/tools/use_mock_search_op.py b/reme_ai/agent/tools/use_mock_search_op.py deleted file mode 100644 index 64098dc3..00000000 --- a/reme_ai/agent/tools/use_mock_search_op.py +++ /dev/null @@ -1,184 +0,0 @@ -"""Tool selection and execution operation for mock search tools. - -This module provides an operation that intelligently selects and executes -the most appropriate mock search tool based on query complexity. -""" - -import asyncio -import datetime -import json - -from flowllm.core.context import C -from flowllm.core.enumeration import Role -from flowllm.core.op import BaseAsyncToolOp -from flowllm.core.schema import Message -from flowllm.core.schema import ToolCall -from flowllm.core.utils import Timer -from loguru import logger - -from reme_ai.agent.tools.mock_search_tools import SearchToolA, SearchToolB, SearchToolC -from reme_ai.schema.memory import ToolCallResult - - -@C.register_op() -class UseMockSearchOp(BaseAsyncToolOp): - """Operation that selects and executes the most appropriate mock search tool. - - This operation uses LLM to intelligently select from available search tools - (SearchToolA, SearchToolB, SearchToolC) based on query characteristics, - then executes the selected tool and records performance metrics. - """ - - file_path: str = __file__ - - def __init__(self, llm: str = "qwen3_30b_instruct", **kwargs): - """Initialize the UseMockSearchOp. - - Args: - llm: LLM model name to use for tool selection - **kwargs: Additional arguments passed to BaseAsyncToolOp - """ - super().__init__(llm=llm, save_answer=True, **kwargs) - - def build_tool_call(self) -> ToolCall: - """Build the tool call schema for the search tool selector. - - Returns: - ToolCall object defining the search tool selector interface - """ - return ToolCall( - **{ - "description": ( - "Intelligently selects and executes the most appropriate search tool " - "based on query complexity. Automatically tracks performance metrics " - "and records tool usage for optimization." - ), - "input_schema": { - "query": { - "type": "string", - "description": "query", - "required": True, - }, - }, - }, - ) - - async def select_tool(self, query: str, tool_ops: list[BaseAsyncToolOp]) -> ToolCall | None: - """Select the most appropriate tool for the given query using LLM. - - Args: - query: The search query to process - tool_ops: List of available tool operations to choose from - - Returns: - Selected ToolCall if a tool was chosen, None otherwise - """ - assistant_message = await self.llm.achat( - messages=[Message(role=Role.USER, content=query)], - tools=[x.tool_call for x in tool_ops], - ) - logger.info(f"assistant_message={assistant_message.model_dump_json()}") - if assistant_message.tool_calls: - return assistant_message.tool_calls[0] - - return None - - async def async_execute(self): - """Execute the tool selection and execution workflow. - - This method selects an appropriate tool, executes it, measures performance, - and creates a ToolCallResult with metrics. - """ - query: str = self.input_dict["query"] - logger.info(f"query={query}") - - tool_ops = [ - SearchToolA(), - SearchToolB(), - SearchToolC(), - ] - - # Step 1: Select the appropriate tool using LLM - tool_call = await self.select_tool(query, tool_ops) - - if tool_call is None: - # No tool selected - error_result = ToolCallResult( - create_time=datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S"), - tool_name="None", - input={"query": query}, - output="No appropriate tool was selected for the query", - token_cost=0, - success=False, - time_cost=0.0, - ) - self.set_output(error_result.model_dump_json()) - return - - # Step 2: Execute the selected tool - selected_op = None - for op in tool_ops: - if op.tool_call.name == tool_call.name: - selected_op = op - break - - if selected_op is None: - # Tool not found (should not happen) - error_result = ToolCallResult( - create_time=datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S"), - tool_name=tool_call.name, - input=tool_call.arguments, - output=f"Tool {tool_call.name} not found in available tools", - token_cost=0, - success=False, - time_cost=0.0, - ) - self.set_output(error_result.model_dump_json()) - return - - # Step 3: Execute the tool with timer - timer = Timer("tool execute") - with timer: - await selected_op.async_call(query=query) - selected_op_output = json.loads(selected_op.output) - content = selected_op_output["content"] - success = selected_op_output["success"] - token_cost = len(content) // 4 # Estimate using a method where every 4 characters constitute one token. - - time_cost = timer.time_cost - - # Create ToolCallResult - tool_call_result = ToolCallResult( - create_time=datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S"), - tool_name=tool_call.name, - input={"query": query}, - output=content, - token_cost=token_cost, - success=success, - time_cost=round(time_cost, 3), - ) - - self.set_output(tool_call_result.model_dump_json()) - - -async def async_main(): - """Main function for testing the UseMockSearchOp with various queries.""" - from reme_ai.main import ReMeApp - - async with ReMeApp(): - test_queries = [ - "What is the capital of France?", - "How does quantum computing work?", - "Analyze the impact of artificial intelligence on global economy, employment, and society", - "When was Python programming language created?", - "Compare different types of renewable energy sources", - ] - - for query in test_queries: - op = UseMockSearchOp() - await op.async_call(query=query) - print(op.output) - - -if __name__ == "__main__": - asyncio.run(async_main()) diff --git a/reme_ai/config/__init__.py b/reme_ai/config/__init__.py deleted file mode 100644 index 68829efa..00000000 --- a/reme_ai/config/__init__.py +++ /dev/null @@ -1,14 +0,0 @@ -"""Configuration module for ReMe. - -This module provides configuration parsing capabilities for the ReMe framework. -It includes: - -- ConfigParser: Configuration parser class that extends PydanticConfigParser - to provide configuration parsing with awareness of the current module location -""" - -from .config_parser import ConfigParser - -__all__ = [ - "ConfigParser", -] diff --git a/reme_ai/config/config_parser.py b/reme_ai/config/config_parser.py deleted file mode 100644 index c21cff0f..00000000 --- a/reme_ai/config/config_parser.py +++ /dev/null @@ -1,24 +0,0 @@ -"""Configuration parser module for ReMe. - -This module provides configuration parsing capabilities for the ReMe framework. -It extends the PydanticConfigParser from FlowLLM to provide configuration parsing -with awareness of the current module location. -""" - -from flowllm.core.utils import PydanticConfigParser - - -class ConfigParser(PydanticConfigParser): - """Configuration parser for ReMe framework. - - Extends PydanticConfigParser to provide configuration parsing capabilities - with awareness of the current module location. Uses the default.yaml - configuration file as the default configuration source. - - Attributes: - current_file: Path to the current file, used for relative config file resolution. - default_config_name: Default configuration file name (without .yaml extension). - """ - - current_file: str = __file__ - default_config_name: str = "default" diff --git a/reme_ai/constants/__init__.py b/reme_ai/constants/__init__.py deleted file mode 100644 index bc1e2d1e..00000000 --- a/reme_ai/constants/__init__.py +++ /dev/null @@ -1,92 +0,0 @@ -"""Constants module for ReMe AI. - -This module provides access to all constants used throughout the application, -including common workflow keys and language-specific constants. -""" - -from . import common_constants -from . import language_constants - -# Export all constants from common_constants -from .common_constants import ( - WORKFLOW_NAME, - RESULT, - MEMORIES, - CHAT_MESSAGES, - CHAT_MESSAGES_SCATTER, - CHAT_KWARGS, - USER_NAME, - TARGET_NAME, - MEMORY_MANAGER, - QUERY_WITH_TS, - RETRIEVE_MEMORY_NODES, - RANKED_MEMORY_NODES, - NOT_REFLECTED_NODES, - NOT_UPDATED_NODES, - EXTRACT_TIME_DICT, - NEW_OBS_NODES, - NEW_OBS_WITH_TIME_NODES, - INSIGHT_NODES, - TODAY_NODES, - MERGE_OBS_NODES, - TIME_INFER, -) - -# Export all constants from language_constants -from .language_constants import ( - DATATIME_WORD_LIST, - WEEKDAYS, - MONTH_DICT, - NONE_WORD, - REPEATED_WORD, - CONTRADICTORY_WORD, - CONTAINED_WORD, - COLON_WORD, - COMMA_WORD, - DEFAULT_HUMAN_NAME, - DATATIME_KEY_MAP, - TIME_INFER_WORD, - USER_NAME_EXPRESSION, -) - -__all__ = [ - # Module exports - "common_constants", - "language_constants", - # Common constants - "WORKFLOW_NAME", - "RESULT", - "MEMORIES", - "CHAT_MESSAGES", - "CHAT_MESSAGES_SCATTER", - "CHAT_KWARGS", - "USER_NAME", - "TARGET_NAME", - "MEMORY_MANAGER", - "QUERY_WITH_TS", - "RETRIEVE_MEMORY_NODES", - "RANKED_MEMORY_NODES", - "NOT_REFLECTED_NODES", - "NOT_UPDATED_NODES", - "EXTRACT_TIME_DICT", - "NEW_OBS_NODES", - "NEW_OBS_WITH_TIME_NODES", - "INSIGHT_NODES", - "TODAY_NODES", - "MERGE_OBS_NODES", - "TIME_INFER", - # Language constants - "DATATIME_WORD_LIST", - "WEEKDAYS", - "MONTH_DICT", - "NONE_WORD", - "REPEATED_WORD", - "CONTRADICTORY_WORD", - "CONTAINED_WORD", - "COLON_WORD", - "COMMA_WORD", - "DEFAULT_HUMAN_NAME", - "DATATIME_KEY_MAP", - "TIME_INFER_WORD", - "USER_NAME_EXPRESSION", -] diff --git a/reme_ai/constants/common_constants.py b/reme_ai/constants/common_constants.py deleted file mode 100644 index aeb7794b..00000000 --- a/reme_ai/constants/common_constants.py +++ /dev/null @@ -1,49 +0,0 @@ -"""Common constants' module. - -This module defines constants used as keys throughout the application to maintain -a consistent reference for data structures related to workflow management, chat -interactions, context storage, memory operations, node processing, and temporal -inference functionalities. -""" - -WORKFLOW_NAME = "workflow_name" - -RESULT = "result" - -MEMORIES = "memories" - -CHAT_MESSAGES = "chat_messages" - -CHAT_MESSAGES_SCATTER = "chat_messages_scatter" - -CHAT_KWARGS = "chat_kwargs" - -USER_NAME = "user_name" - -TARGET_NAME = "target_name" - -MEMORY_MANAGER = "memory_manager" - -QUERY_WITH_TS = "query_with_ts" - -RETRIEVE_MEMORY_NODES = "retrieve_memory_nodes" - -RANKED_MEMORY_NODES = "ranked_memory_nodes" - -NOT_REFLECTED_NODES = "not_reflected_nodes" - -NOT_UPDATED_NODES = "not_updated_nodes" - -EXTRACT_TIME_DICT = "extract_time_dict" - -NEW_OBS_NODES = "new_obs_nodes" - -NEW_OBS_WITH_TIME_NODES = "new_obs_with_time_nodes" - -INSIGHT_NODES = "insight_nodes" - -TODAY_NODES = "today_nodes" - -MERGE_OBS_NODES = "merge_obs_nodes" - -TIME_INFER = "time_infer" diff --git a/reme_ai/constants/language_constants.py b/reme_ai/constants/language_constants.py deleted file mode 100644 index c24e2cc0..00000000 --- a/reme_ai/constants/language_constants.py +++ /dev/null @@ -1,260 +0,0 @@ -"""Language constants module. - -This module provides language-specific constants and mappings for datetime -expressions, weekdays, months, and other linguistic elements used throughout -the application. It supports multiple languages (currently Chinese and English) -and facilitates internationalization of temporal and linguistic processing. -""" - -from ..enumeration.language_enum import LanguageEnum - -# This dictionary maps languages to lists of words related to datetime expressions. -# It aids in recognizing and processing datetime mentions in text, enhancing the system's ability to understand -# temporal context across different languages. -DATATIME_WORD_LIST = { - LanguageEnum.CN: [ - "天", - "周", - "月", - "年", - "星期", - "点", - "分钟", - "小时", - "秒", - "上午", - "下午", - "早上", - "早晨", - "晚上", - "中午", - "日", - "夜", - "清晨", - "傍晚", - "凌晨", - "岁", - ], - LanguageEnum.EN: [ - # Units of Time - "year", - "yr", - "month", - "mo", - "week", - "wk", - "day", - "d", - "hour", - "hr", - "minute", - "min", - "second", - "sec", - # Days of the Week - "Monday", - "Mon", - "Tuesday", - "Tue", - "Tues", - "Wednesday", - "Wed", - "Thursday", - "Thu", - "Thur", - "Thurs", - "Friday", - "Fri", - "Saturday", - "Sat", - "Sunday", - "Sun", - # Months of the Year - "January", - "Jan", - "February", - "Feb", - "March", - "Mar", - "April", - "Apr", - "May", - "May", - "June", - "Jun", - "July", - "Jul", - "August", - "Aug", - "September", - "Sep", - "Sept", - "October", - "Oct", - "November", - "Nov", - "December", - "Dec", - # Relative Time References - "Today", - "Tomorrow", - "Tmrw", - "Yesterday", - "Yday", - "Now", - "Morning", - "AM", - "a.m.", - "Afternoon", - "PM", - "p.m.", - "Evening", - "Night", - "Midnight", - "Noon", - # Seasonal References - "Spring", - "Summer", - "Autumn", - "Fall", - "Winter", - # General Time References - "Century", - "cent.", - "Decade", - "Millennium", - "Quarter", - "Q1", - "Q2", - "Q3", - "Q4", - "Semester", - "Fortnight", - "Weekend", - ], -} - -# A mapping of weekdays for each supported language, facilitating calendar-related operations and understanding -# within the application. -WEEKDAYS = { - LanguageEnum.CN: [ - "周一", - "周二", - "周三", - "周四", - "周五", - "周六", - "周日", - ], - LanguageEnum.EN: [ - "Monday", - "Tuesday", - "Wednesday", - "Thursday", - "Friday", - "Saturday", - "Sunday", - ], -} - -MONTH_DICT = { - LanguageEnum.CN: [ - "1月", - "2月", - "3月", - "4月", - "5月", - "6月", - "7月", - "8月", - "9月", - "10月", - "11月", - "12月", - ], - LanguageEnum.EN: [ - "January", - "February", - "March", - "April", - "May", - "June", - "July", - "August", - "September", - "October", - "November", - "December", - ], -} - -# Constants for the word 'none' in different languages -NONE_WORD = { - LanguageEnum.CN: "无", - LanguageEnum.EN: "none", -} - -# Constants for the word 'repeated' in different languages -REPEATED_WORD = { - LanguageEnum.CN: "重复", - LanguageEnum.EN: "repeated", -} - -# Constants for the word 'contradictory' in different languages -CONTRADICTORY_WORD = { - LanguageEnum.CN: "矛盾", - LanguageEnum.EN: "contradiction", -} - -# Constants for the phrase 'included' in different languages -CONTAINED_WORD = { - LanguageEnum.CN: "被包含", - LanguageEnum.EN: "contained", -} - -# Constants for the symbol ':' in different languages' representations -COLON_WORD = { - LanguageEnum.CN: ":", - LanguageEnum.EN: ":", -} - -# Constants for the symbol ',' in different languages' representations -COMMA_WORD = { - LanguageEnum.CN: ",", - LanguageEnum.EN: ",", -} - -# Default human name placeholders for different languages -DEFAULT_HUMAN_NAME = { - LanguageEnum.CN: "用户", - LanguageEnum.EN: "user", -} - -# Mapping of datetime terms from natural language to standardized keys for each supported language -DATATIME_KEY_MAP = { - LanguageEnum.CN: { - "年": "year", - "月": "month", - "日": "day", - "周": "week", - "星期几": "weekday", - }, - LanguageEnum.EN: { - "Year": "year", - "Month": "month", - "Day": "day", - "Week": "week", - "Weekday": "weekday", - }, -} - -# Phrase for indicating inferred time in different languages -TIME_INFER_WORD = { - LanguageEnum.CN: "推断时间", - LanguageEnum.EN: "Inference time", -} - -USER_NAME_EXPRESSION = { - LanguageEnum.CN: "用户姓名是{name}。", - LanguageEnum.EN: "User's name is {name}.", -} diff --git a/reme_ai/enumeration/__init__.py b/reme_ai/enumeration/__init__.py deleted file mode 100644 index d62d01fd..00000000 --- a/reme_ai/enumeration/__init__.py +++ /dev/null @@ -1,13 +0,0 @@ -"""Enumeration module for ReMe. - -This module provides enumerations used throughout the ReMe system, -including language enumerations and other type definitions. -""" - -from reme_ai.enumeration.working_summary_mode import WorkingSummaryMode -from reme_ai.enumeration.language_enum import LanguageEnum - -__all__ = [ - "WorkingSummaryMode", - "LanguageEnum", -] diff --git a/reme_ai/enumeration/language_enum.py b/reme_ai/enumeration/language_enum.py deleted file mode 100644 index d50ae0f7..00000000 --- a/reme_ai/enumeration/language_enum.py +++ /dev/null @@ -1,20 +0,0 @@ -"""Language enumeration module. - -This module provides enumerations for supported languages in the ReMe system. -""" - -from enum import Enum - - -class LanguageEnum(str, Enum): - """ - An enumeration representing supported languages. - - Members: - - CN: Represents the Chinese language. - - EN: Represents the English language. - """ - - CN = "cn" - - EN = "en" diff --git a/reme_ai/enumeration/working_summary_mode.py b/reme_ai/enumeration/working_summary_mode.py deleted file mode 100644 index a1dde18b..00000000 --- a/reme_ai/enumeration/working_summary_mode.py +++ /dev/null @@ -1,22 +0,0 @@ -"""Working summary mode enumeration module. - -This module defines the strategies for working-memory style summarization in the -ReMe system. -""" - -from enum import Enum - - -class WorkingSummaryMode(str, Enum): - """ - Enumeration representing working summary strategies. - - Members: - - COMPACT: Only compact verbose tool messages into previews. - - COMPRESS: Only apply LLM-based compression over the history. - - AUTO: First compact messages, then optionally compress if needed. - """ - - COMPACT = "compact" - COMPRESS = "compress" - AUTO = "auto" diff --git a/reme_ai/main.py b/reme_ai/main.py deleted file mode 100644 index 74410035..00000000 --- a/reme_ai/main.py +++ /dev/null @@ -1,235 +0,0 @@ -""" -ReMeApp - Reflexive Memory Application - -This module provides the main application class for the ReMe (Reflexive Memory) system, -which extends FlowLLM with specialized memory management capabilities including: -- Task Memory: Store and retrieve task execution histories -- Tool Memory: Track tool usage patterns and experiences -- Personal Memory: Manage user preferences and personal information -""" - -import asyncio -import sys - -from flowllm.core.application import Application -from flowllm.core.context import C -from flowllm.core.schema import FlowResponse - -from reme_ai.config.config_parser import ConfigParser - - -class ReMeApp(Application): - """ - ReMeApp - Main application class for Reflexive Memory system. - - ReMeApp extends FlowLLMApp to provide enhanced memory capabilities for AI agents. - It manages multiple types of memories and provides both synchronous and asynchronous - execution interfaces for memory-enhanced workflows. - """ - - def __init__( - self, - *args, - llm_api_key: str = None, - llm_api_base: str = None, - embedding_api_key: str = None, - embedding_api_base: str = None, - config_path: str = None, - **kwargs, - ): - """ - Initialize ReMeApp with configuration for LLM, embeddings, and vector stores. - - ⚠️ IMPORTANT: The initialization parameters here are consistent with the command-line - startup parameters shown in README.md. You can use the same configuration in both ways: - - Command-line startup: - ```bash - reme \ - backend=http \ - http.port=8002 \ - llm.default.model_name=qwen3-30b-a3b-thinking-2507 \ - embedding_model.default.model_name=text-embedding-v4 \ - vector_store.default.backend=memory - ``` - - Python API equivalent: - ```python - app = ReMeApp( - "llm.default.model_name=qwen3-30b-a3b-thinking-2507", - "embedding_model.default.model_name=text-embedding-v4", - "vector_store.default.backend=memory" - ) - ``` - - Both approaches accept the same configuration parameters and produce identical results. - - Args: - *args: Additional command-line style arguments passed to parser. - These parameters are identical to the command-line startup parameters in README. - - Common configuration examples: - For complete configuration reference, see: reme_ai/config/default.yaml - - LLM Configuration: - - "llm.default.model_name=qwen3-30b-a3b-thinking-2507" - Set LLM model - - "llm.default.backend=openai_compatible" - Set LLM backend type - - "llm.default.params={'temperature': '0.6'}" - Set model parameters - - Embedding Configuration: - - "embedding_model.default.model_name=text-embedding-v4" - Set embedding model - - "embedding_model.default.backend=openai_compatible" - Set embedding backend - - "embedding_model.default.params={'dimensions': 1024}" - Embedding parameters - - Vector Store Configuration: - - "vector_store.default.backend=local" - Use local vector store - - "vector_store.default.backend=memory" - Use memory vector store - - "vector_store.default.backend=qdrant" - Use Qdrant vector store - - "vector_store.default.backend=elasticsearch" - Use Elasticsearch - - "vector_store.default.embedding_model=default" - Link vector store to embedding model - - "vector_store.default.params={'collection_name': 'my_memories'}" - Vector store parameters - llm_api_key: API key for LLM service (e.g., OpenAI, Claude). - If provided, this will override the FLOW_LLM_API_KEY environment variable. - Environment variable: FLOW_LLM_API_KEY - llm_api_base: Base URL for LLM API. Use this for custom or self-hosted endpoints. - If provided, this will override the FLOW_LLM_BASE_URL environment variable. - Example: "https://api.openai.com/v1" - Environment variable: FLOW_LLM_BASE_URL - embedding_api_key: API key for embedding service. Can be different from llm_api_key - if using separate services for embeddings. - If provided, this will override the FLOW_EMBEDDING_API_KEY environment variable. - Environment variable: FLOW_EMBEDDING_API_KEY - embedding_api_base: Base URL for embedding API. For custom embedding endpoints. - If provided, this will override the FLOW_EMBEDDING_BASE_URL environment variable. - Environment variable: FLOW_EMBEDDING_BASE_URL - config_path: Path to custom configuration YAML file. If provided, loads configuration from this file. - Example: "path/to/my_config.yaml" - This overrides the default configuration with your custom settings. - **kwargs: Additional keyword arguments passed to parser. Same format as args but as key-value pairs. - Example: model_name="gpt-4", temperature=0.7 - - Raises: - AssertionError: If required configurations are missing or invalid. - - Note: - - Parameters here mirror the command-line options in README.md exactly - - API keys can be provided via arguments or environment variables (see example.env) - - The parser (ConfigParser) handles merging default configs with custom overrides - - Vector store configuration determines where memories are persisted - - For detailed startup examples and all available parameters, refer to README.md Quick Start section - - See Also: - - README.md "Quick Start" section for command-line startup examples - - README.md "Environment Configuration" for environment variable setup - - example.env for all available environment variables - """ - super().__init__( - *args, - llm_api_key=llm_api_key, - llm_api_base=llm_api_base, - embedding_api_key=embedding_api_key, - embedding_api_base=embedding_api_base, - service_config=None, - parser=ConfigParser, - config_path=config_path, - load_default_config=True, - **kwargs, - ) - - async def async_execute(self, name: str, **kwargs) -> dict: - """ - Asynchronously execute a named flow with given parameters. - - This method executes a registered flow (workflow) by name and returns the result. - Flows are defined in the configuration and registered during app initialization. - - Args: - name: Name of the flow to execute. Must be registered in C.flow_dict. - Common flows in ReMe: - - "task_memory_flow": Query and manage task memories - - "tool_memory_flow": Retrieve tool usage experiences - - "personal_memory_flow": Access personal preferences - - "sop_memory_flow": Execute standard operating procedures - **kwargs: Keyword arguments passed to the flow execution. - Arguments vary by flow type. Common parameters: - - query (str): User query or instruction - - context (dict): Additional context for the flow - - max_results (int): Maximum number of results to return - - threshold (float): Similarity threshold for retrieval - - Returns: - dict: Flow execution result as a dictionary containing: - - status: Execution status (success/failure) - - result: Flow output data - - metadata: Additional execution metadata - - Raises: - AssertionError: If the flow name is not registered in C.flow_dict. - - Example: - ```python - result = await app.async_execute( - "task_memory_flow", - query="Show me all Python debugging tasks", - max_results=10 - ) - print(result['result']) - ``` - """ - assert name in C.flow_dict, f"Invalid flow_name={name} !" - result: FlowResponse = await self.async_execute_flow(name=name, **kwargs) - return result.model_dump() - - def execute(self, name: str, **kwargs) -> dict: - """ - Synchronously execute a named flow with given parameters. - - This is a convenience wrapper around async_execute() for synchronous contexts. - It internally uses asyncio.run() to execute the async flow. - - Args: - name: Name of the flow to execute. See async_execute() for available flows. - **kwargs: Keyword arguments passed to the flow. See async_execute() for details. - - Returns: - dict: Flow execution result. Same format as async_execute(). - - Raises: - AssertionError: If the flow name is not registered. - - Example: - ```python - app = ReMeApp() - result = app.execute( - "tool_memory_flow", - query="How to use the search tool effectively?" - ) - print(result) - ``` - - Note: - For better performance in async contexts, prefer using async_execute() directly. - This method creates a new event loop for each call, which has overhead. - """ - return asyncio.run(self.async_execute(name=name, **kwargs)) - - -def main(): - """ - Entry point for running ReMeApp as a service. - - This function initializes ReMeApp with command-line arguments and starts the service. - It's typically called when running the module directly (python -m reme_ai.app). - - Command-line arguments are passed directly to ReMeApp.__init__(), allowing - configuration via command line: - - Note: - Press Ctrl+C to gracefully shutdown the service. - """ - with ReMeApp(*sys.argv[1:]) as app: - app.run_service() - - -if __name__ == "__main__": - main() diff --git a/reme_ai/retrieve/__init__.py b/reme_ai/retrieve/__init__.py deleted file mode 100644 index 142d2630..00000000 --- a/reme_ai/retrieve/__init__.py +++ /dev/null @@ -1,20 +0,0 @@ -"""Retrieval module for memory operations. - -This module provides submodules for different types of memory retrieval: -- personal: Personal memory retrieval operations -- task: Task memory retrieval operations -- tool: Tool memory retrieval operations -- working: Working-memory retrieval operations -""" - -from . import personal -from . import task -from . import tool -from . import working - -__all__ = [ - "personal", - "task", - "tool", - "working", -] diff --git a/reme_ai/retrieve/personal/__init__.py b/reme_ai/retrieve/personal/__init__.py deleted file mode 100644 index ee4fe03a..00000000 --- a/reme_ai/retrieve/personal/__init__.py +++ /dev/null @@ -1,23 +0,0 @@ -"""Personal memory retrieval operations module. - -This module provides operations for retrieving, ranking, and formatting personal memories -from a vector store, including time extraction, semantic ranking, and memory formatting. -""" - -from .extract_time_op import ExtractTimeOp -from .fuse_rerank_op import FuseRerankOp -from .print_memory_op import PrintMemoryOp -from .read_message_op import ReadMessageOp -from .retrieve_memory_op import RetrieveMemoryOp -from .semantic_rank_op import SemanticRankOp -from .set_query_op import SetQueryOp - -__all__ = [ - "ExtractTimeOp", - "FuseRerankOp", - "PrintMemoryOp", - "ReadMessageOp", - "RetrieveMemoryOp", - "SemanticRankOp", - "SetQueryOp", -] diff --git a/reme_ai/retrieve/personal/extract_time_op.py b/reme_ai/retrieve/personal/extract_time_op.py deleted file mode 100644 index a493f5f0..00000000 --- a/reme_ai/retrieve/personal/extract_time_op.py +++ /dev/null @@ -1,115 +0,0 @@ -"""Time extraction operation for personal memories. - -This module provides functionality to extract time-related information from queries -using LLM-based extraction and pattern matching. -""" - -import re -from typing import Dict - -from flowllm.core.context import C -from flowllm.core.enumeration import Role -from flowllm.core.op import BaseAsyncOp -from flowllm.core.schema import Message -from loguru import logger - -from reme_ai.constants.common_constants import QUERY_WITH_TS, EXTRACT_TIME_DICT -from reme_ai.constants.language_constants import DATATIME_KEY_MAP -from reme_ai.utils.datetime_handler import DatetimeHandler - - -@C.register_op() -class ExtractTimeOp(BaseAsyncOp): - """ - A specialized worker class designed to identify and extract time-related information - from text generated by an LLM, translating date-time keywords based on the set language, - and storing this extracted data within a shared context. - """ - - file_path: str = __file__ - EXTRACT_TIME_PATTERN = r"-\s*(\S+)[::]\s*(\S+)" - - def get_language_value(self, value_dict: dict): - """ - Get value from dictionary based on current language setting. - - Args: - value_dict: Dictionary with language keys - - Returns: - Value for current language or English fallback - """ - return value_dict.get(self.language, value_dict.get("en")) - - async def async_execute(self): - """ - Executes the primary logic of identifying and extracting time data from an LLM's response. - - This method first checks if the input query contains any datetime keywords. If not, it logs and returns. - It then constructs a prompt with contextual information including formatted timestamps and calls the LLM. - The response is parsed for time-related data using regex, translated via a language-specific key map, - and the resulting time data is stored in the shared context. - """ - query, query_timestamp = self.context[QUERY_WITH_TS] - - # Identify if the query contains datetime keywords - contain_datetime = DatetimeHandler.has_time_word(query, self.language) - if not contain_datetime: - logger.info(f"Query contains no datetime keywords: {contain_datetime}") - # Set empty time dict for downstream operations - self.context[EXTRACT_TIME_DICT] = {} - return - - # Prepare the prompt with necessary contextual details - time_format = self.prompt_format(prompt_name="time_string_format") - query_time_str = DatetimeHandler(dt=query_timestamp).string_format(time_format, self.language) - - # Create message with system and few-shot examples - system_prompt = self.prompt_format(prompt_name="extract_time_system") - few_shot = self.prompt_format(prompt_name="extract_time_few_shot") - user_prompt = self.prompt_format( - prompt_name="extract_time_user_query", - query=query, - query_time_str=query_time_str, - ) - - full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_prompt}" - logger.info(f"Extracting time from query: {query[:100]}...") - - # Invoke the LLM to generate a response - response = await self.llm.achat([Message(role=Role.USER, content=full_prompt)]) - - # Handle empty or unsuccessful responses - if not response or not response.content: - logger.warning("LLM returned empty response for time extraction") - self.context[EXTRACT_TIME_DICT] = {} - return - - response_text = response.content - - # Extract and parse time information from the LLM's response - extract_time_dict = self._parse_time_from_response(response_text) - - logger.info(f"Extracted time information: {extract_time_dict}") - self.context[EXTRACT_TIME_DICT] = extract_time_dict - - def _parse_time_from_response(self, response_text: str) -> Dict[str, str]: - """ - Parse time information from LLM response using regex. - - Args: - response_text: Raw LLM response content - - Returns: - Dictionary of extracted time information - """ - extract_time_dict: Dict[str, str] = {} - matches = re.findall(self.EXTRACT_TIME_PATTERN, response_text) - key_map: dict = DATATIME_KEY_MAP[DatetimeHandler.language_transform] - - for key, value in matches: - if key in key_map.keys(): - extract_time_dict[key_map[key]] = value - - logger.debug(f"Time extraction - Response: {response_text[:200]}... Matches: {matches}") - return extract_time_dict diff --git a/reme_ai/retrieve/personal/fuse_rerank_op.py b/reme_ai/retrieve/personal/fuse_rerank_op.py deleted file mode 100644 index 766c66c2..00000000 --- a/reme_ai/retrieve/personal/fuse_rerank_op.py +++ /dev/null @@ -1,228 +0,0 @@ -"""Fuse reranking operation for personal memories. - -This module provides functionality to rerank memory nodes by combining scores, -memory types, and temporal relevance to improve retrieval quality. -""" - -from typing import Dict, List - -from flowllm.core.context import C -from flowllm.core.op import BaseAsyncOp -from loguru import logger - -from reme_ai.constants.common_constants import EXTRACT_TIME_DICT -from reme_ai.schema.memory import BaseMemory - - -@C.register_op() -class FuseRerankOp(BaseAsyncOp): - """ - Reranks the memory nodes by scores, types, and temporal relevance. Formats the top-K reranked nodes to print. - """ - - file_path: str = __file__ - - @staticmethod - def match_memory_time(extract_time_dict: Dict[str, str], memory: BaseMemory): - """ - Determines whether the memory is relevant based on time matching. - - Args: - extract_time_dict: Dictionary containing extracted time information - memory: Memory object to check for time relevance - - Returns: - Tuple of (match_event_flag, match_msg_flag) indicating temporal matches - """ - if extract_time_dict: - match_event_flag = True - for k, v in extract_time_dict.items(): - event_value = memory.metadata.get(f"event_{k}", "") - if event_value in ["-1", v]: - continue - match_event_flag = False - break - - match_msg_flag = True - for k, v in extract_time_dict.items(): - msg_value = memory.metadata.get(f"msg_{k}", "") - if msg_value == v: - continue - match_msg_flag = False - break - else: - match_event_flag = False - match_msg_flag = False - - memory.metadata["match_event_flag"] = str(int(match_event_flag)) - memory.metadata["match_msg_flag"] = str(int(match_msg_flag)) - return match_event_flag, match_msg_flag - - async def async_execute(self): - """ - Executes the reranking process on memories considering their scores, types, and temporal relevance. - - This method performs the following steps: - 1. Retrieves extraction time data and a list of memories from the context. - 2. Reranks memories based on a combination of their original score, type, - and temporal alignment with extracted events/messages. - 3. Selects the top-K reranked memories according to the predefined threshold. - 4. Formats the final list of memories for output. - 5. Sets both response.answer and response.metadata["memory_list"] - """ - # Get operation parameters - fuse_score_threshold = self.op_params.get("fuse_score_threshold", 0.1) - fuse_ratio_dict = self.op_params.get( - "fuse_ratio_dict", - { - "conversation": 0.5, - "observation": 1, - "obs_customized": 1.2, - "insight": 2.0, - }, - ) - fuse_time_ratio = self.op_params.get("fuse_time_ratio", 2.0) - output_memory_max_count = self.op_params.get("output_memory_max_count", 5) - - # Parse input parameters from the context - extract_time_dict: Dict[str, str] = self.context.get(EXTRACT_TIME_DICT, {}) - memory_list: List[BaseMemory] = self.context.response.metadata.get("memory_list", []) - - # Check if memories are available; warn and return if not - if not memory_list: - logger.warning("No memories available for fuse reranking") - self.context.response.answer = "" - self.context.response.metadata["memory_list"] = [] - return - - logger.info(f"Fuse reranking {len(memory_list)} memories with time dict: {bool(extract_time_dict)}") - - # Perform reranking based on score, type, and time relevance - reranked_memories = self._apply_fuse_reranking( - memory_list, - extract_time_dict, - fuse_score_threshold, - fuse_ratio_dict, - fuse_time_ratio, - ) - - # Sort and select top-k memories - reranked_memories = sorted( - reranked_memories, - key=lambda x: x.score or 0.0, - reverse=True, - )[:output_memory_max_count] - - logger.info(f"Final reranked memories: {len(reranked_memories)}") - - # Format memories for output - formatted_memories = self._format_memories_for_output(reranked_memories) - - # Store results in context - both answer and metadata as required - self.context.response.metadata["memory_list"] = reranked_memories - self.context.response.answer = "\n".join(formatted_memories) - - def _apply_fuse_reranking( - self, - memory_list: List[BaseMemory], - extract_time_dict: Dict[str, str], - fuse_score_threshold: float, - fuse_ratio_dict: Dict[str, float], - fuse_time_ratio: float, - ) -> List[BaseMemory]: - """ - Apply fuse reranking logic to memories. - - Args: - memory_list: List of memories to rerank - extract_time_dict: Dictionary containing extracted time information - fuse_score_threshold: Minimum score threshold for memories - fuse_ratio_dict: Dictionary mapping memory types to score multipliers - fuse_time_ratio: Multiplier for time-relevant memories - - Returns: - List of reranked memories with updated scores - """ - reranked_memories = [] - - for memory in memory_list: - # Skip memories below the fuse score threshold - memory_score = memory.score or 0.0 - if memory_score < fuse_score_threshold: - continue - - # Calculate type-based adjustment factor - memory_type = memory.metadata.get("memory_type", "default") - if memory_type not in fuse_ratio_dict: - logger.debug(f"Memory type '{memory_type}' not in fuse_ratio_dict, using default 0.1") - type_ratio: float = fuse_ratio_dict.get(memory_type, 0.1) - - # Determine time relevance adjustment factor - match_event_flag, match_msg_flag = self.match_memory_time(extract_time_dict, memory) - time_ratio: float = fuse_time_ratio if match_event_flag or match_msg_flag else 1.0 - - # Apply reranking score adjustments - original_score = memory_score - memory.score = memory_score * type_ratio * time_ratio - - logger.debug( - f"Memory reranked: {original_score:.3f} -> {memory.score:.3f} " - f"(type={type_ratio}, time={time_ratio})", - ) - - reranked_memories.append(memory) - - return reranked_memories - - def _format_memories_for_output(self, memories: List[BaseMemory]) -> List[str]: - """ - Format memories for final output. - - Args: - memories: List of memories to format - - Returns: - List of formatted memory strings - """ - formatted_memories = [] - - for memory in memories: - # Log reranking details - logger.info( - f"Final memory: Score={memory.score:.3f}, " - f"Event={memory.metadata.get('match_event_flag', '0')}, " - f"Msg={memory.metadata.get('match_msg_flag', '0')}, " - f"Content={memory.content[:50]}...", - ) - - # Format memory with timestamp if available - formatted_content = self._format_memory_with_timestamp(memory, self.language) - formatted_memories.append(formatted_content) - - return formatted_memories - - @staticmethod - def _format_memory_with_timestamp(memory, language: str = "en") -> str: - """ - Format memory content with timestamp if available. - - Args: - memory: Memory object - language: Language for formatting - - Returns: - Formatted memory content string - """ - try: - if hasattr(memory, "timestamp") and memory.timestamp: - from reme_ai.utils.datetime_handler import DatetimeHandler - - dt_handler = DatetimeHandler(memory.timestamp) - datetime_str = dt_handler.datetime_format("%Y-%m-%d %H:%M:%S") - weekday = dt_handler.get_dt_info_dict(language)["weekday"] - return f"[{datetime_str} {weekday}] {memory.content}" - else: - return memory.content - except Exception as e: - logger.warning(f"Failed to format memory with timestamp: {e}") - return memory.content diff --git a/reme_ai/retrieve/personal/print_memory_op.py b/reme_ai/retrieve/personal/print_memory_op.py deleted file mode 100644 index 8cef06bd..00000000 --- a/reme_ai/retrieve/personal/print_memory_op.py +++ /dev/null @@ -1,149 +0,0 @@ -"""Memory printing operation for personal memories. - -This module provides functionality to format and print memories in various formats -for display or output purposes. -""" - -from typing import List - -from flowllm.core.context import C -from flowllm.core.op import BaseAsyncOp -from loguru import logger - -from reme_ai.schema.memory import BaseMemory - - -@C.register_op() -class PrintMemoryOp(BaseAsyncOp): - """ - Formats the memories to print. - """ - - file_path: str = __file__ - - async def async_execute(self): - """ - Executes the primary function, it involves: - 1. Fetches the memories. - 2. Formats them for printing. - 3. Set the formatted string back into the context - """ - # Get memory list from context - memory_list: List[BaseMemory] = self.context.response.metadata.get("memory_list", []) - - if not memory_list: - logger.info("No memories to print") - self.context.response.answer = "No memories found." - return - - logger.info(f"Formatting {len(memory_list)} memories for printing") - - # Format memories for printing - formatted_memories = self._format_memories_for_print(memory_list) - - # Store result in context - self.context.response.answer = formatted_memories - logger.info(f"Formatted memories: {formatted_memories}") - - @staticmethod - def _format_memories_for_print(memories: List[BaseMemory]) -> str: - """ - Format memories for printing. - - Args: - memories: List of memory objects to format - - Returns: - Formatted string representation of memories - """ - if not memories: - return "No memories available." - - formatted_memories = [] - - for i, memory in enumerate(memories, 1): - memory_text = f"Memory {i}:\n" - memory_text += f" When to use: {memory.when_to_use}\n" - memory_text += f" Content: {memory.content}\n" - - # Add additional metadata if available - if hasattr(memory, "metadata") and memory.metadata: - metadata_items = [] - for key, value in memory.metadata.items(): - if key not in ["when_to_use", "content"]: - metadata_items.append(f"{key}: {value}") - if metadata_items: - memory_text += f" Metadata: {', '.join(metadata_items)}\n" - - formatted_memories.append(memory_text) - - return "\n".join(formatted_memories) - - @staticmethod - def format_memories_for_output(memories: List) -> str: - """ - Format memory list for output string. - - Args: - memories: List of memory objects - - Returns: - Formatted string - """ - if not memories: - return "" - - formatted_parts = [] - for i, memory in enumerate(memories, 1): - when_to_use = getattr(memory, "when_to_use", "") or memory.get("when_to_use", "") - content = getattr(memory, "content", "") or memory.get("content", "") - - part = f"Memory {i}:\n" - if when_to_use: - part += f"When to use: {when_to_use}\n" - if content: - part += f"Content: {content}\n" - - formatted_parts.append(part) - - return "\n".join(formatted_parts) - - @staticmethod - def format_memories_for_simple_output(memories: List) -> str: - """ - Format memory list for simple flow output. - - Args: - memories: List of memory objects - - Returns: - Formatted string suitable for response.answer - """ - if not memories: - return "No relevant memories found." - - content_parts = ["Previous Memory"] - - for memory in memories: - # Safely get field values - when_to_use = getattr(memory, "when_to_use", "") or memory.get("when_to_use", "") - content = getattr(memory, "content", "") or memory.get("content", "") - - # Skip memories with empty content - if not content: - continue - - # Format individual memory - memory_text = f"- when_to_use: {when_to_use}\n content: {content}" - content_parts.append(memory_text) - - # If no valid memories, return empty message - if len(content_parts) == 1: # Only title - return "No relevant memories with valid content found." - - content_parts.append( - "\nPlease consider the helpful parts from these in answering the question, " - "to make the response more comprehensive and substantial.", - ) - - return "\n".join(content_parts) diff --git a/reme_ai/retrieve/personal/read_message_op.py b/reme_ai/retrieve/personal/read_message_op.py deleted file mode 100644 index 1f901f42..00000000 --- a/reme_ai/retrieve/personal/read_message_op.py +++ /dev/null @@ -1,60 +0,0 @@ -"""Message reading operation for personal memories. - -This module provides functionality to read and filter unmemorized chat messages -from the context for processing. -""" - -from typing import List - -from flowllm.core.context import C -from flowllm.core.op import BaseAsyncOp -from flowllm.core.schema import Message -from loguru import logger - - -@C.register_op() -class ReadMessageOp(BaseAsyncOp): - """ - Fetches unmemorized chat messages. - """ - - file_path: str = __file__ - - async def async_execute(self): - """ - Executes the primary function to fetch unmemorized chat messages. - """ - # Get chat messages from context - chat_messages = self.context.chat_messages - target_name = self.context.target_name - contextual_msg_max_count = self.op_params.get("contextual_msg_max_count", 10) - - chat_messages_not_memorized: List[List[Message]] = [] - for messages in chat_messages: - if not messages: - continue - - if hasattr(messages[0], "memorized") and messages[0].memorized: - continue - - contain_flag = False - - for msg in messages: - if hasattr(msg, "role_name") and msg.role_name == target_name: - contain_flag = True - break - - if contain_flag: - chat_messages_not_memorized.append(messages) - - chat_message_scatter = [] - for messages in chat_messages_not_memorized[-contextual_msg_max_count:]: - chat_message_scatter.extend(messages) - - # Sort by time_created if available - if chat_message_scatter and hasattr(chat_message_scatter[0], "time_created"): - chat_message_scatter.sort(key=lambda _: _.time_created) - - # Store result in context - self.context.chat_messages = chat_message_scatter - logger.info(f"Retrieved {len(chat_message_scatter)} unmemorized chat messages") diff --git a/reme_ai/retrieve/personal/retrieve_memory_op.py b/reme_ai/retrieve/personal/retrieve_memory_op.py deleted file mode 100644 index e0d0fb2b..00000000 --- a/reme_ai/retrieve/personal/retrieve_memory_op.py +++ /dev/null @@ -1,61 +0,0 @@ -"""Memory retrieval operation for personal memories. - -This module provides functionality to retrieve memories from a vector store -based on query similarity and score thresholds. -""" - -from typing import List - -from flowllm.core.context import C -from flowllm.core.op import BaseAsyncOp -from flowllm.core.schema import VectorNode -from loguru import logger - -from reme_ai.schema.memory import BaseMemory, vector_node_to_memory - - -@C.register_op() -class RetrieveMemoryOp(BaseAsyncOp): - """ - Retrieves memories based on specified criteria such as status, type, and timestamp. - Processes these memories concurrently, sorts them by similarity, and logs the activity, - facilitating efficient memory retrieval operations within a given scope. - """ - - async def async_execute(self): - """ - Executes the memory retrieval operation. - - This method: - 1. Retrieves memories from vector store based on query similarity - 2. Removes duplicate memories based on content - 3. Filters memories by score threshold if specified - 4. Stores the retrieved memories in context metadata - """ - recall_key: str = self.op_params.get("recall_key", "query") - top_k: int = self.context.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] = await self.vector_store.async_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 diff --git a/reme_ai/retrieve/personal/semantic_rank_op.py b/reme_ai/retrieve/personal/semantic_rank_op.py deleted file mode 100644 index 2b1a0b88..00000000 --- a/reme_ai/retrieve/personal/semantic_rank_op.py +++ /dev/null @@ -1,211 +0,0 @@ -"""Semantic ranking operation for personal memories. - -This module provides functionality to rank memories semantically using LLM-based -relevance scoring to improve retrieval quality. -""" - -import json -import re -from typing import List - -from flowllm.core.context import C -from flowllm.core.enumeration import Role -from flowllm.core.op import BaseAsyncOp -from flowllm.core.schema import Message -from loguru import logger - -from reme_ai.schema.memory import BaseMemory - - -@C.register_op() -class SemanticRankOp(BaseAsyncOp): - """ - The SemanticRankOp class processes queries by retrieving memory nodes, - removing duplicates, ranking them based on semantic relevance using a model, - assigning scores, sorting the nodes, and storing the ranked nodes back, - while logging relevant information. - """ - - file_path: str = __file__ - - async def async_execute(self): - """ - Executes the primary workflow of the SemanticRankOp which includes: - - Retrieves query and memory list from context. - - Removes duplicate memories. - - Ranks memories semantically using LLM. - - Assigns scores to memories. - - Sorts memories by score. - - Saves the ranked memories back to context. - - If no memories are retrieved or if the ranking fails, - appropriate warnings are logged. - """ - # Get memory list from context - previous op guarantees this exists - memory_list: List[BaseMemory] = self.context.response.metadata["memory_list"] - query: str = self.context.query - - # Get parameters from op_params - enable_ranker: bool = self.op_params.get("enable_ranker", True) - output_memory_max_count: int = self.op_params.get("output_memory_max_count", 10) - - if not memory_list: - logger.warning("Memory list is empty!") - return - - logger.info(f"Semantic ranking {len(memory_list)} memories for query: {query[:100]}...") - - if not enable_ranker or len(memory_list) <= output_memory_max_count: - # Use original scores if ranker is disabled or memory count is small - logger.info("Skipping semantic ranking - using original scores") - else: - # Remove duplicates based on content - memory_dict = {memory.content.strip(): memory for memory in memory_list if memory.content.strip()} - memory_list = list(memory_dict.values()) - logger.info(f"After deduplication: {len(memory_list)} memories") - - # Perform semantic ranking using LLM - ranked_memories = await self._semantic_rank_memories(query, memory_list) - if ranked_memories: - memory_list = ranked_memories - - # Sort by score - memory_list = sorted(memory_list, key=lambda m: m.score, reverse=True) - - # Log top ranked memories - logger.info(f"Semantic ranking completed for query: {query[:50]}...") - for i, memory in enumerate(memory_list[:5]): # Log top 5 - logger.info(f"Top {i + 1}: Score={memory.score:.3f}, Content={memory.content[:80]}...") - - # Save ranked memories back to context - self.context.response.metadata["memory_list"] = memory_list - - 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. - - Args: - query: User query to rank memories against - memories: List of memories to rank - - Returns: - List of memories with updated semantic scores - """ - if not memories: - return memories - - # Format memories for ranking - formatted_memories = SemanticRankOp.format_memories_for_llm_ranking(memories) - - # Create prompt for semantic ranking - prompt = f"""Given the query: "{query}" - -Please rank the following memories by their semantic relevance to the query. -Rate each memory on a scale of 0.0 to 1.0 where 1.0 is most relevant. - -Memories: -{formatted_memories} - -Please respond in JSON format: -{{"rankings": [{{"index": 0, "score": 0.8}}, {{"index": 1, "score": 0.6}}, ...]}}""" - - 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") - return memories - - # Parse and apply ranking results - rankings = SemanticRankOp.parse_llm_ranking_response(response.content) - - if rankings: - applied_count = SemanticRankOp.apply_semantic_scores_to_memories(memories, rankings) - logger.info(f"Successfully applied semantic rankings to {applied_count} memories") - else: - logger.warning("Failed to parse ranking results") - - return memories - - @staticmethod - def parse_llm_ranking_response(response: str) -> List[dict]: - """ - Parse LLM ranking response to extract rankings. - - Args: - response: Raw LLM response string containing ranking JSON - - Returns: - List of ranking dictionaries with index and score - """ - try: - # Try to extract JSON blocks - json_pattern = r"```json\s*([\s\S]*?)\s*```" - json_blocks = re.findall(json_pattern, response) - - if json_blocks: - parsed = json.loads(json_blocks[0]) - if isinstance(parsed, dict) and "rankings" in parsed: - return parsed["rankings"] - - # Fallback: try to parse the entire response as JSON - parsed = json.loads(response) - if isinstance(parsed, dict) and "rankings" in parsed: - return parsed["rankings"] - - except json.JSONDecodeError: - logger.warning("Failed to parse ranking response as JSON") - - return [] - - @staticmethod - def apply_semantic_scores_to_memories(memories: List, rankings: List[dict]) -> int: - """ - Apply semantic ranking scores to memory objects. - - Args: - memories: List of memory objects to update - rankings: List of ranking dictionaries with index and score - - Returns: - Number of memories successfully updated with scores - """ - applied_count = 0 - - for ranking in rankings: - idx = ranking.get("index", -1) - score = ranking.get("score", 0.0) - - if 0 <= idx < len(memories): - # Set score on memory object - if hasattr(memories[idx], "score"): - memories[idx].score = score - applied_count += 1 - else: - # Add score as metadata if score attribute doesn't exist - if not hasattr(memories[idx], "metadata"): - memories[idx].metadata = {} - memories[idx].metadata["semantic_score"] = score - applied_count += 1 - - return applied_count - - @staticmethod - def format_memories_for_llm_ranking(memories: List) -> str: - """ - Format memories for LLM ranking input. - - Args: - memories: List of memory objects to format - - Returns: - Formatted string representation of memories for LLM input - """ - formatted_memories = [] - - for i, memory in enumerate(memories): - memory_text = f"Memory {i}:\n" - memory_text += f"When to use: {memory.when_to_use}\n" - memory_text += f"Content: {memory.content}\n" - formatted_memories.append(memory_text) - - return "\n---\n".join(formatted_memories) diff --git a/reme_ai/retrieve/personal/set_query_op.py b/reme_ai/retrieve/personal/set_query_op.py deleted file mode 100644 index f382f4a5..00000000 --- a/reme_ai/retrieve/personal/set_query_op.py +++ /dev/null @@ -1,44 +0,0 @@ -"""Query setting operation for personal memories. - -This module provides functionality to set query and timestamp in the context -for downstream memory retrieval operations. -""" - -import datetime -from typing import Tuple - -from flowllm.core.context import C -from flowllm.core.op import BaseAsyncOp -from loguru import logger - -from reme_ai.constants.common_constants import QUERY_WITH_TS - - -@C.register_op() -class SetQueryOp(BaseAsyncOp): - """ - The `SetQueryOp` class is responsible for setting a query and its associated timestamp - into the context, utilizing either provided parameters or details from the context. - """ - - 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. - - Input requirement: self.context.query must exist (flow input requirement) - """ - # Flow guarantees query exists - use it directly - query: str = self.context.query - timestamp: int = int(datetime.datetime.now().timestamp()) - - # Set timestamp if provided in op_params - _timestamp = self.op_params.get("timestamp") - if _timestamp and isinstance(_timestamp, int): - timestamp = _timestamp - - # Store the query and its timestamp in the context - query_with_ts: Tuple[str, int] = (query, timestamp) - self.context[QUERY_WITH_TS] = query_with_ts - - logger.info(f"Set query with timestamp: query='{query}', timestamp={timestamp}") diff --git a/reme_ai/retrieve/task/__init__.py b/reme_ai/retrieve/task/__init__.py deleted file mode 100644 index fb4e3dd7..00000000 --- a/reme_ai/retrieve/task/__init__.py +++ /dev/null @@ -1,17 +0,0 @@ -"""Task memory retrieval operations module. - -This module provides operations for building queries, reranking memories, -rewriting memory context, and merging memories for task-related retrieval. -""" - -from .build_query_op import BuildQueryOp -from .merge_memory_op import MergeMemoryOp -from .rerank_memory_op import RerankMemoryOp -from .rewrite_memory_op import RewriteMemoryOp - -__all__ = [ - "BuildQueryOp", - "MergeMemoryOp", - "RerankMemoryOp", - "RewriteMemoryOp", -] diff --git a/reme_ai/retrieve/task/build_query_op.py b/reme_ai/retrieve/task/build_query_op.py deleted file mode 100644 index d0339c6f..00000000 --- a/reme_ai/retrieve/task/build_query_op.py +++ /dev/null @@ -1,61 +0,0 @@ -"""Query building operation module. - -This module provides functionality to build retrieval queries from either -explicit query strings or conversation messages, optionally using LLM to -generate optimized queries. -""" - -from flowllm.core.context import C -from flowllm.core.enumeration import Role -from flowllm.core.op import BaseAsyncOp -from flowllm.core.schema import Message -from flowllm.core.utils import merge_messages_content -from loguru import logger - - -@C.register_op() -class BuildQueryOp(BaseAsyncOp): - """Build retrieval query from context or messages. - - This operation constructs a query string for memory retrieval. It can use - an explicit query from context, or generate one from conversation messages - using either LLM-based generation or simple message concatenation. - """ - - file_path: str = __file__ - - async def async_execute(self): - """Execute the query building operation. - - Builds a query string from either: - 1. An explicit query in the context - 2. Conversation messages (using LLM or simple concatenation) - - Stores the built query in context.query. - """ - if "query" in self.context: - query = self.context.query - - elif "messages" in self.context: - if self.op_params.get("enable_llm_build", True): - execution_process = merge_messages_content(self.context.messages) - prompt = self.prompt_format(prompt_name="query_build", execution_process=execution_process) - message = await self.llm.achat(messages=[Message(role=Role.USER, content=prompt)]) - query = message.content - - else: - context_parts = [] - message_summaries = [] - for message in self.context.messages[-3:]: # Last 3 messages - content = message.content[:200] + "..." if len(message.content) > 200 else message.content - message_summaries.append(f"- {message.role.value}: {content}") - if message_summaries: - context_parts.append("Recent messages:\n" + "\n".join(message_summaries)) - - query = "\n\n".join(context_parts) - - else: - raise RuntimeError("query or messages is required!") - - logger.info(f"build.query={query}") - self.context.query = query diff --git a/reme_ai/retrieve/task/merge_memory_op.py b/reme_ai/retrieve/task/merge_memory_op.py deleted file mode 100644 index c56e4106..00000000 --- a/reme_ai/retrieve/task/merge_memory_op.py +++ /dev/null @@ -1,47 +0,0 @@ -"""Memory merging operation module. - -This module provides functionality to merge multiple retrieved memories -into a single formatted context string for use in LLM responses. -""" - -from typing import List - -from flowllm.core.context import C -from flowllm.core.op import BaseAsyncOp -from loguru import logger - -from reme_ai.schema.memory import BaseMemory - - -@C.register_op() -class MergeMemoryOp(BaseAsyncOp): - """Merge multiple memories into a single formatted context. - - This operation takes a list of retrieved memories and formats them into - a single context string that can be used to guide LLM responses. It includes - instructions for the LLM to consider the helpful parts from these memories. - """ - - async def async_execute(self): - """Execute the memory merging operation. - - Merges memories from context metadata into a formatted string with - instructions for the LLM. Stores the merged result in response.answer. - """ - memory_list: List[BaseMemory] = self.context.response.metadata["memory_list"] - - if not memory_list: - return - - content_collector = ["Previous Memory"] - for memory in memory_list: - if not memory.content: - continue - - content_collector.append(f"- {memory.when_to_use} {memory.content}\n") - content_collector.append( - "Please consider the helpful parts from these in answering the question, " - "to make the response more comprehensive and substantial.", - ) - self.context.response.answer = "\n".join(content_collector) - logger.info(f"response.answer={self.context.response.answer}") diff --git a/reme_ai/retrieve/task/rerank_memory_op.py b/reme_ai/retrieve/task/rerank_memory_op.py deleted file mode 100644 index beb464e3..00000000 --- a/reme_ai/retrieve/task/rerank_memory_op.py +++ /dev/null @@ -1,197 +0,0 @@ -"""Memory reranking operation module. - -This module provides functionality to rerank and filter retrieved memories -using LLM-based reranking and score-based filtering to select the most relevant -memories for the current task. -""" - -import json -import re -from typing import List - -from flowllm.core.context import C -from flowllm.core.enumeration import Role -from flowllm.core.op import BaseAsyncOp -from flowllm.core.schema import Message -from loguru import logger - -from reme_ai.schema.memory import BaseMemory - - -@C.register_op() -class RerankMemoryOp(BaseAsyncOp): - """Rerank and filter recalled experiences using LLM and score-based filtering. - - This operation takes recalled memories and applies multiple filtering and - ranking strategies to select the most relevant memories for the current task. - It supports LLM-based reranking and score-based filtering. - """ - - file_path: str = __file__ - - async def async_execute(self): - """Execute the memory reranking operation. - - Applies LLM-based reranking (optional) and score-based filtering (optional) - to select the top-k most relevant memories. Stores the reranked results - in the context response metadata. - """ - memory_list: List[BaseMemory] = self.context.response.metadata["memory_list"] - retrieval_query: str = self.context.query - enable_llm_rerank = self.op_params.get("enable_llm_rerank", True) - enable_score_filter = self.op_params.get("enable_score_filter", False) - min_score_threshold = self.op_params.get("min_score_threshold", 0.3) - top_k = self.op_params.get("top_k", 5) - - logger.info(f"top_k: {top_k}") - - if not memory_list: - logger.info("No recalled memory_list to rerank") - return - - logger.info(f"Reranking {len(memory_list)} memories") - - # Step 1: LLM reranking (optional) - if enable_llm_rerank: - memory_list = await self._llm_rerank(retrieval_query, memory_list) - logger.info(f"After LLM reranking: {len(memory_list)} memories") - - # Step 2: Score-based filtering (optional) - if enable_score_filter: - memory_list = self._score_based_filter(memory_list, min_score_threshold) - logger.info(f"After score filtering: {len(memory_list)} memories") - - # Step 3: Return top-k results - reranked_memories = memory_list[:top_k] - logger.info(f"Final reranked results: {len(reranked_memories)} memories") - - # Store results in context - self.context.response.metadata["memory_list"] = reranked_memories - - async def _llm_rerank(self, query: str, candidates: List[BaseMemory]) -> List[BaseMemory]: - """LLM-based reranking of candidate experiences. - - Args: - query: The retrieval query used to rank candidates. - candidates: List of memory candidates to rerank. - - Returns: - List of memories reranked by relevance to the query. - """ - if not candidates: - return candidates - - # Format candidates for LLM evaluation - candidates_text = self._format_candidates_for_rerank(candidates) - - prompt = self.prompt_format( - prompt_name="memory_rerank_prompt", - query=query, - candidates=candidates_text, - num_candidates=len(candidates), - ) - - response = await self.llm.achat([Message(role=Role.USER, content=prompt)]) - - # Parse reranking results - reranked_indices = self._parse_rerank_response(response.content) - - # Reorder candidates based on LLM ranking - if reranked_indices: - reranked_candidates = [] - for idx in reranked_indices: - if 0 <= idx < len(candidates): - reranked_candidates.append(candidates[idx]) - - # Add any remaining candidates that weren't explicitly ranked - ranked_indices_set = set(reranked_indices) - for i, candidate in enumerate(candidates): - if i not in ranked_indices_set: - reranked_candidates.append(candidate) - - return reranked_candidates - - return candidates - - @staticmethod - def _score_based_filter(memories: List[BaseMemory], min_score: float) -> List[BaseMemory]: - """Filter memories based on quality scores. - - Args: - memories: List of memories to filter. - min_score: Minimum combined score threshold for filtering. - - Returns: - List of memories that meet the minimum score threshold. - """ - filtered_memories = [] - - for memory in memories: - # Get confidence score from metadata - confidence = memory.metadata.get("confidence", 0.5) - validation_score = memory.score or 0.5 - - # Calculate combined score - combined_score = (confidence + validation_score) / 2 - - if combined_score >= min_score: - filtered_memories.append(memory) - else: - logger.debug(f"Filtered out memory with score {combined_score:.2f}") - - logger.info(f"Score filtering: {len(filtered_memories)}/{len(memories)} memories retained") - return filtered_memories - - @staticmethod - def _format_candidates_for_rerank(candidates: List[BaseMemory]) -> str: - """Format candidates for LLM reranking. - - Args: - candidates: List of memory candidates to format. - - Returns: - Formatted string representation of candidates for LLM evaluation. - """ - formatted_candidates = [] - - for i, candidate in enumerate(candidates): - condition = candidate.when_to_use - content = candidate.content - - candidate_text = f"Candidate {i}:\n" - candidate_text += f"Condition: {condition}\n" - candidate_text += f"Experience: {content}\n" - - formatted_candidates.append(candidate_text) - - return "\n---\n".join(formatted_candidates) - - @staticmethod - def _parse_rerank_response(response: str) -> List[int]: - """Parse LLM reranking response to extract ranked indices. - - Args: - response: The LLM response containing ranked indices. - - Returns: - List of indices representing the reranked order. - """ - try: - # Try to extract JSON format - json_pattern = r"```json\s*([\s\S]*?)\s*```" - json_blocks = re.findall(json_pattern, response) - - if json_blocks: - parsed = json.loads(json_blocks[0]) - if isinstance(parsed, dict) and "ranked_indices" in parsed: - return parsed["ranked_indices"] - elif isinstance(parsed, list): - return parsed - - # Try to extract numbers from text - numbers = re.findall(r"\b\d+\b", response) - return [int(num) for num in numbers if int(num) < 100] # Reasonable upper bound - - except Exception as e: - logger.error(f"Error parsing rerank response: {e}") - return [] diff --git a/reme_ai/retrieve/task/rewrite_memory_op.py b/reme_ai/retrieve/task/rewrite_memory_op.py deleted file mode 100644 index 708fbf0e..00000000 --- a/reme_ai/retrieve/task/rewrite_memory_op.py +++ /dev/null @@ -1,205 +0,0 @@ -"""Memory rewriting operation module. - -This module provides functionality to rewrite and format retrieved memories -into context messages that can be used by LLMs for task completion. -""" - -import json -import re -from typing import List - -from flowllm.core.context import C -from flowllm.core.enumeration import Role -from flowllm.core.op import BaseAsyncOp -from flowllm.core.schema import Message -from loguru import logger - -from reme_ai.schema.memory import BaseMemory - - -@C.register_op() -class RewriteMemoryOp(BaseAsyncOp): - """Generate and rewrite context messages from reranked experiences. - - This operation takes reranked memories and formats them into context messages - that can be used by LLMs. It optionally uses LLM-based rewriting to make - the context more relevant and actionable for the current task. - """ - - file_path: str = __file__ - - async def async_execute(self): - """Execute the memory rewrite operation. - - Retrieves memories from context metadata, formats them, and optionally - rewrites them using LLM to make them more relevant for the current query. - Stores the rewritten context in the response answer field. - """ - memory_list: List[BaseMemory] = self.context.response.metadata["memory_list"] - query: str = self.context.query - messages: List[Message] = [Message(**x) if isinstance(x, dict) else x for x in self.context.get("messages", [])] - - if not memory_list: - logger.info("No reranked memories to rewrite") - self.context.response.answer = "" - return - - logger.info(f"Generating context from {len(memory_list)} memories") - - # Generate initial context message - rewritten_memory = await self._generate_context_message(query, messages, memory_list) - - # Store results in context - self.context.response.answer = rewritten_memory - self.context.response.metadata["memory_list"] = [memory.model_dump() for memory in memory_list] - - async def _generate_context_message(self, query: str, messages: List[Message], memories: List[BaseMemory]) -> str: - """Generate context message from retrieved memories. - - Args: - query: The current query string. - messages: List of conversation messages for context. - memories: List of retrieved memories to format. - - Returns: - Formatted context string, optionally rewritten by LLM. - """ - if not memories: - return "" - - try: - logger.info("memories") - # Format retrieved memories - formatted_memories = self._format_memories_for_context(memories) - - if self.op_params.get("enable_llm_rewrite", True): - context_content = await self._rewrite_context(query, formatted_memories, messages) - else: - context_content = formatted_memories - - return context_content - - except Exception as e: - logger.error(f"Error generating context message: {e}") - return self._format_memories_for_context(memories) - - async def _rewrite_context(self, query: str, context_content: str, messages: List[Message]) -> str: - """LLM-based context rewriting to make experiences more relevant and actionable. - - Args: - query: The current query string. - context_content: The formatted context content to rewrite. - messages: List of conversation messages for additional context. - - Returns: - Rewritten context string optimized for the current task. - """ - if not context_content: - return context_content - - try: - # Extract current context - current_context = self._extract_context(messages) - - prompt = self.prompt_format( - prompt_name="memory_rewrite_prompt", - current_query=query, - current_context=current_context, - original_context=context_content, - ) - - response = await self.llm.achat([Message(role=Role.USER, content=prompt)]) - - # Extract rewritten context - rewritten_context = self._parse_json_response(response.content, "rewritten_context") - - if rewritten_context and rewritten_context.strip(): - logger.info("Context successfully rewritten for current task") - return rewritten_context.strip() - - return context_content - - except Exception as e: - logger.error(f"Error in context rewriting: {e}") - return context_content - - @staticmethod - def _format_memories_for_context(memories: List[BaseMemory]) -> str: - """Format memories for context generation. - - Args: - memories: List of memories to format. - - Returns: - Formatted string containing all memories with their conditions and content. - """ - formatted_memories = [] - - for i, memory in enumerate(memories, 1): - condition = memory.when_to_use - memory_content = memory.content - memory_text = f"Memory {i} :\n When to use: {condition}\n Content: {memory_content}\n" - - formatted_memories.append(memory_text) - - return "\n".join(formatted_memories) - - @staticmethod - def _extract_context(messages: List[Message]) -> str: - """Extract relevant context from messages. - - Args: - messages: List of conversation messages. - - Returns: - Formatted string containing recent conversation context. - """ - if not messages: - return "" - - context_parts = [] - - # Add recent messages if available - recent_messages = messages[-3:] # Last 3 messages - message_summaries = [] - for message in recent_messages: - content = message.content[:300] + "..." if len(message.content) > 300 else message.content - message_summaries.append(f"- {message.role.value}: {content}") - - if message_summaries: - context_parts.append("Recent conversation:\n" + "\n".join(message_summaries)) - - return "\n\n".join(context_parts) - - @staticmethod - def _parse_json_response(response: str, key: str) -> str: - """Parse JSON response to extract specific key. - - Args: - response: The response string that may contain JSON. - key: The key to extract from the JSON object. - - Returns: - The value associated with the key, or the response string if parsing fails. - """ - try: - # Try to extract JSON blocks - json_pattern = r"```json\s*([\s\S]*?)\s*```" - json_blocks = re.findall(json_pattern, response) - - if json_blocks: - parsed = json.loads(json_blocks[0]) - if isinstance(parsed, dict) and key in parsed: - return parsed[key] - - # Fallback: try to parse the entire response as JSON - parsed = json.loads(response) - if isinstance(parsed, dict) and key in parsed: - return parsed[key] - - except json.JSONDecodeError: - logger.warning(f"Failed to parse JSON response for key '{key}', using raw response") - # If JSON parsing fails, return the response as-is for fallback - return response.strip() - - return "" diff --git a/reme_ai/retrieve/tool/__init__.py b/reme_ai/retrieve/tool/__init__.py deleted file mode 100644 index 05af6e6e..00000000 --- a/reme_ai/retrieve/tool/__init__.py +++ /dev/null @@ -1,11 +0,0 @@ -"""Tool memory retrieval operations module. - -This module provides operations for retrieving tool memories from a vector store -based on tool names, including formatting and matching tool memory results. -""" - -from .retrieve_tool_memory_op import RetrieveToolMemoryOp - -__all__ = [ - "RetrieveToolMemoryOp", -] diff --git a/reme_ai/retrieve/tool/retrieve_tool_memory_op.py b/reme_ai/retrieve/tool/retrieve_tool_memory_op.py deleted file mode 100644 index 2c27728e..00000000 --- a/reme_ai/retrieve/tool/retrieve_tool_memory_op.py +++ /dev/null @@ -1,125 +0,0 @@ -"""Tool memory retrieval operation module. - -This module provides functionality to retrieve tool memories from a vector store -based on tool names, format them into structured documents, and match them with -the requested tools. -""" - -from typing import List - -from flowllm.core.context import C -from flowllm.core.op import BaseAsyncOp -from flowllm.core.schema import VectorNode -from loguru import logger - -from reme_ai.schema.memory import ToolMemory, vector_node_to_memory - - -@C.register_op() -class RetrieveToolMemoryOp(BaseAsyncOp): - """Retrieves tool memories from vector store based on tool names. - - This operation searches for tool memories in the vector store using tool names, - validates that the retrieved memories match the requested tools, and formats - them into a structured document format for use in the context. - """ - - file_path: str = __file__ - - @staticmethod - def _format_tool_memories(memories: List[ToolMemory]) -> str: - """Format tool memories into a structured document format. - - Args: - memories: List of ToolMemory objects to format. - - Returns: - A formatted string containing all tool memories with separators. - """ - lines = [f"Retrieved {len(memories)} tool memory(ies):\n"] - - for idx, memory in enumerate(memories, 1): - lines.append(f"Tool: {memory.when_to_use}") - lines.append(memory.content) - - if idx < len(memories): - lines.append("\n---\n") - - return "\n".join(lines) - - async def async_execute(self): - """Execute the tool memory retrieval operation. - - This method: - 1. Extracts tool names from context - 2. Searches for each tool in the vector store - 3. Validates that retrieved memories match the requested tools - 4. Formats the memories into a structured document - 5. Stores the results in context response - - The operation expects 'tool_names' in the context, which should be a - comma-separated string of tool names. For each tool name, it retrieves - the top matching memory from the vector store and validates that it - matches exactly. - """ - tool_names: str = self.context.get("tool_names", "") - workspace_id: str = self.context.workspace_id - - if not tool_names: - logger.warning("tool_names is empty, skipping processing") - self.context.response.answer = "tool_names is required" - self.context.response.success = False - return - - # Split tool names by comma - tool_name_list = [name.strip() for name in tool_names.split(",") if name.strip()] - logger.info(f"workspace_id={workspace_id} retrieving {len(tool_name_list)} tools: {tool_name_list}") - - # Search for each tool in the vector store - matched_tool_memories: List[ToolMemory] = [] - - for tool_name in tool_name_list: - nodes: List[VectorNode] = await self.vector_store.async_search( - query=tool_name, - workspace_id=workspace_id, - top_k=1, - ) - - if nodes: - top_node = nodes[0] - memory = vector_node_to_memory(top_node) - - # Ensure it's a ToolMemory and when_to_use matches - if isinstance(memory, ToolMemory) and memory.when_to_use == tool_name: - matched_tool_memories.append(memory) - logger.info( - f"Found tool_memory for tool_name={tool_name}, " - f"memory_id={memory.memory_id}, " - f"total_calls={len(memory.tool_call_results)}", - ) - else: - logger.warning(f"No exact match found for tool_name={tool_name}") - else: - logger.warning(f"No memory found for tool_name={tool_name}") - - if not matched_tool_memories: - logger.info("No matching tool memories found") - self.context.response.answer = "No matching tool memories found" - self.context.response.success = False - return - - # Format tool memories as document - formatted_answer = self._format_tool_memories(matched_tool_memories) - - # Set response - self.context.response.answer = formatted_answer - self.context.response.success = True - self.context.response.metadata["memory_list"] = matched_tool_memories - - # Log retrieval results - for memory in matched_tool_memories: - logger.info( - f"Retrieved tool: {memory.when_to_use}, " - f"total_calls={len(memory.tool_call_results)}, " - f"content_length={len(memory.content)}", - ) diff --git a/reme_ai/retrieve/working/__init__.py b/reme_ai/retrieve/working/__init__.py deleted file mode 100644 index ab1181a8..00000000 --- a/reme_ai/retrieve/working/__init__.py +++ /dev/null @@ -1,19 +0,0 @@ -"""Working operations package for ReMe retrieve framework. - -This package provides working/file-related operations that can be used in LLM-powered -flows. It currently includes ready-to-use operations for: - -- BatchWriteFileOp: Batch write multiple files operation -- GrepOp: Text search operation for finding patterns in files -- ReadFileOp: Read single file operation -""" - -from .batch_write_file_op import BatchWriteFileOp -from .grep_op import GrepOp -from .read_file_op import ReadFileOp - -__all__ = [ - "BatchWriteFileOp", - "GrepOp", - "ReadFileOp", -] diff --git a/reme_ai/retrieve/working/batch_write_file_op.py b/reme_ai/retrieve/working/batch_write_file_op.py deleted file mode 100644 index 44bc9425..00000000 --- a/reme_ai/retrieve/working/batch_write_file_op.py +++ /dev/null @@ -1,44 +0,0 @@ -"""Batch write file operation module. - -This module provides a tool operation for batch writing multiple files at once. -It processes a dictionary of file paths and contents, writing each file sequentially -and returning a combined result of all write operations. -""" - -from flowllm.core.context import C -from flowllm.core.op import BaseAsyncOp -from loguru import logger - -from .write_file_op import WriteFileOp - - -@C.register_op() -class BatchWriteFileOp(BaseAsyncOp): - """Batch write file operation. - - This operation writes multiple files in a single batch. It takes a dictionary - of file paths and contents from the context, and writes each file using - WriteFileOp. Returns a combined result of all write operations. - """ - - def __init__(self, save_answer: bool = False, **kwargs): - super().__init__(**kwargs) - self.save_answer: bool = save_answer - - async def async_execute(self): - """Execute the batch write file operation. - - Reads write_file_dict from context, which should be a dictionary mapping - file paths to file contents. Writes each file sequentially and collects - the results. - """ - # Get write file dictionary from context - write_file_dict: dict = self.context.get("write_file_dict", {}) - if not write_file_dict: - logger.info("No write file task.") - return - - # Process each file in the dictionary - for file_path, content in write_file_dict.items(): - write_op = WriteFileOp(save_answer=self.save_answer) - await write_op.async_call(file_path=file_path, content=content) diff --git a/reme_ai/retrieve/working/grep_op.py b/reme_ai/retrieve/working/grep_op.py deleted file mode 100644 index eaa30642..00000000 --- a/reme_ai/retrieve/working/grep_op.py +++ /dev/null @@ -1,95 +0,0 @@ -"""Grep text search operation module. - -This module provides a tool operation for searching text patterns in files. -It enables efficient content-based search using regular expressions, with support -for glob pattern filtering and result limiting. -""" - -import re -from pathlib import Path - -from flowllm.core.context import C -from flowllm.core.op import BaseAsyncToolOp -from flowllm.core.schema import ToolCall -from loguru import logger - - -@C.register_op() -class GrepOp(BaseAsyncToolOp): - """Grep text search operation. - - This operation searches for text patterns in files using regular expressions. - Supports glob pattern filtering and result limiting. - """ - - file_path = __file__ - - def __init__(self, **kwargs): - kwargs.setdefault("raise_exception", False) - super().__init__(**kwargs) - - def build_tool_call(self) -> ToolCall: - """Build and return the tool call schema for this operator.""" - tool_params = { - "name": "Grep", - "description": self.get_prompt("tool_desc"), - "input_schema": { - "file_path": { - "type": "string", - "description": self.get_prompt("file_path"), - "required": True, - }, - "pattern": { - "type": "string", - "description": self.get_prompt("pattern"), - "required": True, - }, - "limit": { - "type": "number", - "description": self.get_prompt("limit"), - "required": False, - }, - }, - } - - return ToolCall(**tool_params) - - async def async_execute(self): - """Execute the grep search operation.""" - pattern: str = self.input_dict.get("pattern", "").strip() - file_path: str | None | Path = self.input_dict.get("file_path", "") - limit: int = int(self.input_dict.get("limit", 50)) - - assert pattern, "The 'pattern' parameter cannot be empty." - assert file_path, "The 'file_path' parameter is required." - target_file = Path(file_path).expanduser().resolve() - assert target_file.exists(), f"File does not exist: {target_file}" - assert target_file.is_file(), f"Path is not a file: {target_file}" - - logger.info(f"Searching for pattern '{pattern}' in {target_file}") - - regex = re.compile(re.escape(pattern), re.IGNORECASE) - results = [] - - with target_file.open("r", encoding="utf-8", errors="ignore") as f: - for line_num, line in enumerate(f, 1): - if regex.search(line): - results.append(f"{target_file}:{line_num}:{line.rstrip()}") - if len(results) >= limit: - break - - if not results: - search_location = f'in file_path "{file_path}"' if file_path else "in the workspace directory" - result_msg = f'No matches found for pattern "{pattern}" {search_location}.' - else: - result_msg = "\n".join(results) - - self.set_output(result_msg) - - async def async_default_execute(self, e: Exception = None, **_kwargs): - """Fill outputs with a default failure message when execution fails.""" - pattern: str = self.input_dict.get("pattern", "").strip() - error_msg = f'Failed to search for pattern "{pattern}"' - if e: - error_msg += f": {str(e)}" - self.set_output(error_msg) diff --git a/reme_ai/retrieve/working/read_file_op.py b/reme_ai/retrieve/working/read_file_op.py deleted file mode 100644 index 0d2a8f93..00000000 --- a/reme_ai/retrieve/working/read_file_op.py +++ /dev/null @@ -1,91 +0,0 @@ -"""Read file operation module. - -This module provides a tool operation for reading file contents. -It supports reading entire files or specific line ranges for large files. -""" - -from pathlib import Path -from typing import Optional - -from flowllm.core.context import C -from flowllm.core.op import BaseAsyncToolOp -from flowllm.core.schema import ToolCall -from reme_ai.utils.op_utils import run_shell_command - - -@C.register_op() -class ReadFileOp(BaseAsyncToolOp): - """Read file operation. - - This operation reads and returns the content of a specified file. - For text files, it can read specific line ranges using offset and limit. - """ - - file_path = __file__ - - def __init__(self, **kwargs): - kwargs.setdefault("raise_exception", False) - super().__init__(**kwargs) - - def build_tool_call(self) -> ToolCall: - """Build and return the tool call schema for this operator.""" - tool_params = { - "name": "ReadFile", - "description": self.get_prompt("tool_desc"), - "input_schema": { - "file_path": { - "type": "string", - "description": self.get_prompt("file_path"), - "required": True, - }, - "offset": { - "type": "number", - "description": self.get_prompt("offset"), - "required": True, - }, - "limit": { - "type": "number", - "description": self.get_prompt("limit"), - "required": True, - }, - }, - } - - return ToolCall(**tool_params) - - async def async_execute(self): - """Execute the read file operation.""" - file_path: str = self.input_dict.get("file_path", "").strip() - offset: Optional[int] = int(self.input_dict.get("offset")) - limit: Optional[int] = int(self.input_dict.get("limit")) - - # Validate and resolve file path - assert file_path, "The 'file_path' parameter cannot be empty." - file_path_obj = Path(file_path).expanduser().resolve() - assert file_path_obj.exists(), f"File not found: {file_path_obj}" - assert file_path_obj.is_file(), f"Path is not a file: {file_path_obj}" - - # Set default values and validate - offset = offset or 0 - limit = limit or 1000000 - assert offset >= 0, "Offset must be a non-negative number" - assert limit > 0, "Limit must be a positive number" - - # Use sed for efficient line reading (1-indexed) - start_line = offset + 1 - end_line = offset + limit - - cmd = ["sed", "-n", f"{start_line},{end_line}p", str(file_path_obj)] - stdout, stderr, returncode = await run_shell_command(cmd, timeout=30) - - assert returncode == 0, f"sed command failed: {stderr}" - content = stdout.rstrip("\n") - self.set_output(content) - - async def async_default_execute(self, e: Exception = None, **_kwargs): - """Fill outputs with a default failure message when execution fails.""" - file_path: str = self.input_dict.get("file_path", "").strip() - error_msg = f'Failed to read file "{file_path}"' - if e: - error_msg += f": {str(e)}" - self.set_output(error_msg) diff --git a/reme_ai/retrieve/working/write_file_op.py b/reme_ai/retrieve/working/write_file_op.py deleted file mode 100644 index 20097c79..00000000 --- a/reme_ai/retrieve/working/write_file_op.py +++ /dev/null @@ -1,89 +0,0 @@ -"""Write file operation module. - -This module provides a tool operation for writing content to files. -It supports creating new files or overwriting existing files, and automatically -creates parent directories if they don't exist. -""" - -from pathlib import Path - -from flowllm.core.context import C -from flowllm.core.op import BaseAsyncToolOp -from flowllm.core.schema import ToolCall - - -@C.register_op() -class WriteFileOp(BaseAsyncToolOp): - """Write file operation. - - This operation writes content to a specified file. If the file doesn't exist, - it will be created. If parent directories don't exist, they will be created automatically. - """ - - file_path = __file__ - - def __init__(self, **kwargs): - kwargs.setdefault("raise_exception", False) - super().__init__(**kwargs) - - def build_tool_call(self) -> ToolCall: - """Build and return the tool call schema for this operator.""" - tool_params = { - "name": "WriteFile", - "description": self.get_prompt("tool_desc"), - "input_schema": { - "file_path": { - "type": "string", - "description": self.get_prompt("file_path"), - "required": True, - }, - "content": { - "type": "string", - "description": self.get_prompt("content"), - "required": True, - }, - }, - } - - return ToolCall(**tool_params) - - async def async_execute(self): - """Execute the write file operation.""" - file_path: str = self.input_dict.get("file_path", "").strip() - content: str = self.input_dict.get("content", "") - - # Validate file_path - if not file_path: - raise ValueError("The 'file_path' parameter cannot be empty.") - - # Resolve file path - file_path_obj = Path(file_path).expanduser().resolve() - - # Check if path is a directory - if file_path_obj.exists() and file_path_obj.is_dir(): - raise ValueError(f"Path is a directory, not a file: {file_path_obj}") - - # Create parent directories if they don't exist - file_path_obj.parent.mkdir(parents=True, exist_ok=True) - - # Check if file exists - file_exists = file_path_obj.exists() and file_path_obj.is_file() - - # Write content to file - file_path_obj.write_text(content, encoding="utf-8") - - # Format success message - if file_exists: - result = f"Successfully overwrote file: {file_path_obj}" - else: - result = f"Successfully created and wrote to new file: {file_path_obj}" - - self.set_output(result) - - async def async_default_execute(self, e: Exception = None, **_kwargs): - """Fill outputs with a default failure message when execution fails.""" - file_path: str = self.input_dict.get("file_path", "").strip() - error_msg = f'Failed to write file "{file_path}"' - if e: - error_msg += f": {str(e)}" - self.set_output(error_msg) diff --git a/reme_ai/schema/__init__.py b/reme_ai/schema/__init__.py deleted file mode 100644 index 70255638..00000000 --- a/reme_ai/schema/__init__.py +++ /dev/null @@ -1,34 +0,0 @@ -"""Schema module for ReMe. - -This module provides data structures and schemas for memory management, -including memory types, tool call results, and conversion utilities. -""" - -from flowllm.core.enumeration import Role # noqa -from flowllm.core.schema import Message, Trajectory # noqa - -from reme_ai.schema.memory import ( - BaseMemory, - PersonalMemory, - TaskMemory, - ToolCallResult, - ToolMemory, - dict_to_memory, - vector_node_to_memory, -) - -__all__ = [ - # FlowLLM schema imports - "Message", - "Role", - "Trajectory", - # Memory classes - "BaseMemory", - "TaskMemory", - "PersonalMemory", - "ToolMemory", - "ToolCallResult", - # Utility functions - "vector_node_to_memory", - "dict_to_memory", -] diff --git a/reme_ai/schema/memory.py b/reme_ai/schema/memory.py deleted file mode 100644 index 296e0b27..00000000 --- a/reme_ai/schema/memory.py +++ /dev/null @@ -1,597 +0,0 @@ -"""Memory schema definitions for ReMe. - -This module defines the core memory data structures used in the ReMe system, -including base memory classes and specialized memory types for tasks, personal -information, and tool call results. -""" - -import datetime -import hashlib -import json -from abc import ABC -from typing import List -from uuid import uuid4 - -from flowllm.core.schema import VectorNode -from mcp.types import CallToolResult, TextContent -from pydantic import BaseModel, Field - - -class BaseMemory(BaseModel, ABC): - """Base class for all memory types in the ReMe system. - - This abstract base class provides common fields and methods for all memory - types, including workspace identification, content storage, timestamps, - and conversion to/from vector nodes for storage and retrieval. - - Attributes: - workspace_id: Identifier for the workspace this memory belongs to. - memory_id: Unique identifier for this memory instance. - memory_type: Type of memory (task, personal, tool, etc.). - when_to_use: Description of when this memory should be retrieved. - content: The actual content of the memory (string or bytes). - score: Relevance score for this memory (0.0 to 1.0). - time_created: Timestamp when the memory was created. - time_modified: Timestamp when the memory was last modified. - author: Identifier of the entity that created this memory. - metadata: Additional metadata dictionary for extensibility. - """ - - workspace_id: str = Field(default="") - memory_id: str = Field(default_factory=lambda: uuid4().hex) - memory_type: str = Field(default=...) - - when_to_use: str = Field(default="") - content: str | bytes = Field(default="") - score: float = Field(default=0) - - time_created: str = Field(default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")) - time_modified: str = Field(default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")) - author: str = Field(default="") - - metadata: dict = Field(default_factory=dict) - - def update_modified_time(self): - """Update the time_modified field to the current timestamp.""" - self.time_modified = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") - - def update_metadata(self, new_metadata): - """Update the metadata dictionary with new values. - - Args: - new_metadata: Dictionary containing new metadata to replace existing metadata. - """ - self.metadata = new_metadata - - def to_vector_node(self) -> VectorNode: - """Convert this memory instance to a VectorNode for storage. - - Returns: - VectorNode: A vector node representation of this memory. - - Raises: - NotImplementedError: Must be implemented by subclasses. - """ - raise NotImplementedError - - @classmethod - def from_vector_node(cls, node: VectorNode): - """Create a memory instance from a VectorNode. - - Args: - node: VectorNode containing memory data. - - Returns: - BaseMemory: A memory instance reconstructed from the vector node. - - Raises: - NotImplementedError: Must be implemented by subclasses. - """ - raise NotImplementedError - - -class TaskMemory(BaseMemory): - """Memory type for storing task-related information. - - TaskMemory is used to store information about tasks, including when to use - the memory and the task content itself. It extends BaseMemory with - task-specific behavior. - - Attributes: - memory_type: Always set to "task" for task memories. - """ - - memory_type: str = Field(default="task") - - def to_vector_node(self) -> VectorNode: - """Convert this TaskMemory to a VectorNode. - - Returns: - VectorNode: Vector node representation with when_to_use as content - and all other fields stored in metadata. - """ - return VectorNode( - unique_id=self.memory_id, - workspace_id=self.workspace_id, - content=self.when_to_use, - metadata={ - "memory_type": self.memory_type, - "content": self.content, - "score": self.score, - "time_created": self.time_created, - "time_modified": self.time_modified, - "author": self.author, - "metadata": json.dumps(self.metadata, ensure_ascii=False), - }, - ) - - @classmethod - def from_vector_node(cls, node: VectorNode) -> "TaskMemory": - """Create a TaskMemory instance from a VectorNode. - - Args: - node: VectorNode containing task memory data. - - Returns: - TaskMemory: Reconstructed TaskMemory instance. - """ - metadata = node.metadata.copy() - memory_metadata = metadata.pop("metadata", {}) - if isinstance(memory_metadata, str): - memory_metadata = json.loads(memory_metadata) - - return cls( - workspace_id=node.workspace_id, - memory_id=node.unique_id, - memory_type=metadata.pop("memory_type"), - when_to_use=node.content, - content=metadata.pop("content"), - score=metadata.pop("score"), - time_created=metadata.pop("time_created"), - time_modified=metadata.pop("time_modified"), - author=metadata.pop("author"), - metadata=memory_metadata, - ) - - -class PersonalMemory(BaseMemory): - """Memory type for storing personal information and user preferences. - - PersonalMemory extends BaseMemory with fields specific to personal data, - including target information and reflection subject attributes. This is - used for storing user preferences, personal insights, and reflection data. - - Attributes: - memory_type: Always set to "personal" for personal memories. - target: Target identifier or category for this personal memory. - reflection_subject: Subject of reflection for storing reflection attributes. - """ - - memory_type: str = Field(default="personal") - target: str = Field(default="") - reflection_subject: str = Field(default="") # For storing reflection subject attributes - - def to_vector_node(self) -> VectorNode: - """Convert this PersonalMemory to a VectorNode. - - Returns: - VectorNode: Vector node representation with when_to_use as content - and all other fields including target and reflection_subject - stored in metadata. - """ - return VectorNode( - unique_id=self.memory_id, - workspace_id=self.workspace_id, - content=self.when_to_use, - metadata={ - "memory_type": self.memory_type, - "content": self.content, - "target": self.target, - "reflection_subject": self.reflection_subject, - "score": self.score, - "time_created": self.time_created, - "time_modified": self.time_modified, - "author": self.author, - "metadata": json.dumps(self.metadata, ensure_ascii=False), - }, - ) - - @classmethod - def from_vector_node(cls, node: VectorNode) -> "PersonalMemory": - """Create a PersonalMemory instance from a VectorNode. - - Args: - node: VectorNode containing personal memory data. - - Returns: - PersonalMemory: Reconstructed PersonalMemory instance. - """ - metadata = node.metadata.copy() - memory_metadata = metadata.pop("metadata", {}) - if isinstance(memory_metadata, str): - memory_metadata = json.loads(memory_metadata) - - return cls( - workspace_id=node.workspace_id, - memory_id=node.unique_id, - memory_type=metadata.pop("memory_type"), - when_to_use=node.content, - content=metadata.pop("content"), - target=metadata.pop("target", ""), - reflection_subject=metadata.pop("reflection_subject", ""), - score=metadata.pop("score"), - time_created=metadata.pop("time_created"), - time_modified=metadata.pop("time_modified"), - author=metadata.pop("author"), - metadata=memory_metadata, - ) - - -class ToolCallResult(BaseModel): - """Represents the result of a tool invocation. - - This class stores comprehensive information about a tool call, including - inputs, outputs, performance metrics, evaluation, and deduplication hash. - - Attributes: - create_time: Timestamp when the tool was invoked. - tool_name: Name of the tool that was called. - input: Input parameters passed to the tool (dict or string). - output: Output result from the tool execution. - token_cost: Number of tokens consumed by the tool call (-1 if unknown). - success: Whether the tool invocation completed successfully. - time_cost: Time taken for the tool invocation in seconds. - summary: Brief summary of the tool call result. - evaluation: Detailed evaluation of the tool invocation. - score: Quality score from 0.0 (failure) to 1.0 (complete success). - is_summarized: Whether this tool call has been included in a summary. - call_hash: MD5 hash of input and output for deduplication. - metadata: Additional metadata dictionary. - """ - - create_time: str = Field(default="", description="Time of tool invocation") - tool_name: str = Field(default=..., description="Name of the tool") - input: dict | str = Field(default="", description="Tool input") - output: str = Field(default="", description="Tool output") - token_cost: int = Field(default=-1, description="Token consumption of the tool") - success: bool = Field(default=True, description="Whether the tool invocation was successful") - time_cost: float = Field(default=0, description="Time consumed by the tool invocation, in seconds") - summary: str = Field(default="", description="Brief summary of the tool call result") - evaluation: str = Field(default="", description="Detailed evaluation for the tool invocation") - score: float = Field(default=0, description="Score of the Evaluation (0.0 for failure, 1.0 for complete success)") - is_summarized: bool = Field(default=False, description="Whether this tool call has been included in a summary") - call_hash: str = Field(default="", description="Hash value of input and output combined for deduplication") - - metadata: dict = Field(default_factory=dict) - - def generate_hash(self) -> str: - """Generate hash value from tool input and output for deduplication. - - Creates an MD5 hash from the combined input and output strings. - This hash is used to identify duplicate tool calls. - - Returns: - str: MD5 hash hexdigest of the combined input and output. - """ - # Convert input to string if it's a dict - input_str = json.dumps(self.input, sort_keys=True) if isinstance(self.input, dict) else str(self.input) - - # Combine input and output - combined = f"{input_str}|{self.output}" - - # Generate MD5 hash - hash_value = hashlib.md5(combined.encode("utf-8")).hexdigest() - - return hash_value - - def ensure_hash(self): - """Ensure call_hash is set, generate if empty.""" - if not self.call_hash: - self.call_hash = self.generate_hash() - - def from_mcp_tool_result(self, tool_result: CallToolResult, max_char_len: int = None): - """Populate this instance from an MCP CallToolResult. - - Args: - tool_result: MCP CallToolResult to extract data from. - max_char_len: Optional maximum character length for output content. - If provided, output will be truncated to this length. - """ - text_list = [] - for content in tool_result.content: - if isinstance(content, TextContent): - text_list.append(content.text) - - else: - raise NotImplementedError(f"content.type={type(content)} not supported") - content = "\n".join(text_list) - - if max_char_len: - content = content[:max_char_len] - self.output = content - - self.success = not tool_result.is_error - self.metadata.update(tool_result.meta) - - -class ToolMemory(BaseMemory): - """Memory type for storing tool call execution history. - - ToolMemory extends BaseMemory to store a collection of tool call results, - allowing tracking of tool usage patterns, performance metrics, and - execution history for analysis and summarization. - - Attributes: - memory_type: Always set to "tool" for tool memories. - tool_call_results: List of ToolCallResult instances representing - historical tool invocations. - """ - - memory_type: str = Field(default="tool") - tool_call_results: List[ToolCallResult] = Field(default_factory=list) - - def to_vector_node(self) -> VectorNode: - """Convert this ToolMemory to a VectorNode. - - Returns: - VectorNode: Vector node representation with when_to_use as content - and all tool_call_results serialized in metadata. - """ - return VectorNode( - unique_id=self.memory_id, - workspace_id=self.workspace_id, - content=self.when_to_use, - metadata={ - "memory_type": self.memory_type, - "content": self.content, - "score": self.score, - "time_created": self.time_created, - "time_modified": self.time_modified, - "author": self.author, - "tool_call_results": [x.model_dump() for x in self.tool_call_results], - "metadata": json.dumps(self.metadata, ensure_ascii=False), - }, - ) - - def statistic(self, recent_frequency: int = 20) -> dict: - """Calculate statistical information for the most recent N tool calls. - - Analyzes the most recent tool calls and computes average metrics including - token cost, success rate, time cost, and quality scores. - - Args: - recent_frequency: Number of most recent tool calls to analyze. - Defaults to 20. - - Returns: - dict: Dictionary containing: - - avg_token_cost: Average token consumption (rounded to 2 decimals) - - avg_time_cost: Average execution time in seconds (rounded to 3 decimals) - - success_rate: Ratio of successful calls (rounded to 4 decimals) - - avg_score: Average quality score (rounded to 3 decimals) - """ - if not self.tool_call_results: - return { - "total_calls": 0, - "recent_calls_analyzed": 0, - "avg_token_cost": 0.0, - "success_rate": 0.0, - "avg_time_cost": 0.0, - "avg_score": 0.0, - } - - # Get the most recent N tool calls (or all if less than N) - recent_calls = self.tool_call_results[-recent_frequency:] - # total_calls = len(self.tool_call_results) - recent_calls_count = len(recent_calls) - - # Calculate statistics - total_token_cost = sum(call.token_cost for call in recent_calls if call.token_cost >= 0) - valid_token_calls = [call for call in recent_calls if call.token_cost >= 0] - avg_token_cost = total_token_cost / len(valid_token_calls) if valid_token_calls else 0.0 - - successful_calls = sum(1 for call in recent_calls if call.success) - success_rate = successful_calls / recent_calls_count if recent_calls_count > 0 else 0.0 - - total_time_cost = sum(call.time_cost for call in recent_calls) - avg_time_cost = total_time_cost / recent_calls_count if recent_calls_count > 0 else 0.0 - - total_score = sum(call.score for call in recent_calls) - avg_score = total_score / recent_calls_count if recent_calls_count > 0 else 0.0 - - return { - "avg_token_cost": round(avg_token_cost, 2), - "avg_time_cost": round(avg_time_cost, 3), - "success_rate": round(success_rate, 4), - "avg_score": round(avg_score, 3), - } - - @classmethod - def from_vector_node(cls, node: VectorNode) -> "ToolMemory": - """Create a ToolMemory instance from a VectorNode. - - Args: - node: VectorNode containing tool memory data. - - Returns: - ToolMemory: Reconstructed ToolMemory instance with tool_call_results - deserialized from metadata. - """ - metadata = node.metadata.copy() - tool_call_results = [ToolCallResult(**result) for result in metadata.pop("tool_call_results", [])] - memory_metadata = metadata.pop("metadata", {}) - if isinstance(memory_metadata, str): - memory_metadata = json.loads(memory_metadata) - - return cls( - workspace_id=node.workspace_id, - memory_id=node.unique_id, - when_to_use=node.content, - memory_type=metadata.pop("memory_type"), - content=metadata.pop("content"), - score=metadata.pop("score"), - time_created=metadata.pop("time_created"), - time_modified=metadata.pop("time_modified"), - author=metadata.pop("author"), - tool_call_results=tool_call_results, - metadata=memory_metadata, - ) - - -def vector_node_to_memory(node: VectorNode): - """Convert a VectorNode to the appropriate memory type. - - This function inspects the memory_type in the node's metadata and - reconstructs the appropriate memory subclass (TaskMemory, PersonalMemory, - or ToolMemory). - - Args: - node: VectorNode containing memory data with memory_type in metadata. - - Returns: - BaseMemory: Instance of the appropriate memory subclass based on - memory_type. - - Raises: - RuntimeError: If memory_type is not recognized or not present. - """ - memory_type = node.metadata.get("memory_type") - if memory_type == "task": - return TaskMemory.from_vector_node(node) - - elif memory_type == "personal": - return PersonalMemory.from_vector_node(node) - - elif memory_type == "tool": - return ToolMemory.from_vector_node(node) - - else: - raise RuntimeError(f"memory_type={memory_type} not supported!") - - -def dict_to_memory(memory_dict: dict): - """Create a memory instance from a dictionary. - - This function creates the appropriate memory subclass based on the - memory_type field in the dictionary. Defaults to TaskMemory if - memory_type is not specified. - - Args: - memory_dict: Dictionary containing memory data with optional - memory_type field. - - Returns: - BaseMemory: Instance of the appropriate memory subclass based on - memory_type. - - Raises: - RuntimeError: If memory_type is not recognized. - """ - memory_type = memory_dict.get("memory_type", "task") - if memory_type == "task": - return TaskMemory(**memory_dict) - - elif memory_type == "personal": - return PersonalMemory(**memory_dict) - - elif memory_type == "tool": - return ToolMemory(**memory_dict) - - else: - raise RuntimeError(f"memory_type={memory_type} not supported!") - - -def task_main(): - """Test function for TaskMemory serialization and deserialization.""" - e1 = TaskMemory( - workspace_id="w_1024", - memory_id="123", - when_to_use="test case use", - content="test content", - score=0.99, - metadata={}, - ) - print(e1.model_dump_json(indent=2)) - v1 = e1.to_vector_node() - print(v1.model_dump_json(indent=2)) - e2 = vector_node_to_memory(v1) - print(e2.model_dump_json(indent=2)) - - -def personal_main(): - """Test function for PersonalMemory serialization and deserialization.""" - p1 = PersonalMemory( - workspace_id="w_2048", - memory_id="456", - when_to_use="personal memory test case", - content="personal test content", - target="user_preferences", - reflection_subject="learning_style", - score=0.85, - metadata={"category": "user_profile"}, - ) - print("PersonalMemory test:") - print(p1.model_dump_json(indent=2)) - v1 = p1.to_vector_node() - print("VectorNode:") - print(v1.model_dump_json(indent=2)) - p2 = vector_node_to_memory(v1) - print("Reconstructed PersonalMemory:") - print(p2.model_dump_json(indent=2)) - - -def tool_main(): - """Test function for ToolMemory serialization and deserialization.""" - # Create sample tool call results - tool_result1 = ToolCallResult( - create_time="2025-10-15 10:30:00", - tool_name="file_reader", - input={"file_path": "/test/file.txt"}, - output="File content successfully read", - token_cost=50, - success=True, - time_cost=0.5, - evaluation="Successfully executed", - score=0.95, - ) - - tool_result2 = ToolCallResult( - create_time="2025-10-15 10:31:00", - tool_name="data_processor", - input={"data": "sample_data", "format": "json"}, - output="Data processed successfully", - token_cost=75, - success=True, - time_cost=1.2, - evaluation="Good performance", - score=0.88, - ) - - t1 = ToolMemory( - workspace_id="w_4096", - memory_id="789", - memory_type="tool", - when_to_use="tool execution memory test", - content="tool execution test content", - score=0.92, - tool_call_results=[tool_result1, tool_result2], - metadata={"execution_context": "test_environment"}, - ) - - print("ToolMemory test:") - print(t1.model_dump_json(indent=2)) - v1 = t1.to_vector_node() - print("VectorNode:") - print(v1.model_dump_json(indent=2)) - t2 = ToolMemory.from_vector_node(v1) - print("Reconstructed ToolMemory:") - print(t2.model_dump_json(indent=2)) - - -if __name__ == "__main__": - print("=== Task Memory Test ===") - # task_main() - print("\n=== Personal Memory Test ===") - # personal_main() - print("\n=== Tool Memory Test ===") - tool_main() diff --git a/reme_ai/service/__init__.py b/reme_ai/service/__init__.py deleted file mode 100644 index 64eea5bb..00000000 --- a/reme_ai/service/__init__.py +++ /dev/null @@ -1,15 +0,0 @@ -"""Memory service modules for ReMe. - -This package provides memory service implementations for managing different -types of memories including task memories and personal memories. -""" - -from reme_ai.service.agentscope_runtime_memory_service import AgentscopeRuntimeMemoryService -from reme_ai.service.personal_memory_service import PersonalMemoryService -from reme_ai.service.task_memory_service import TaskMemoryService - -__all__ = [ - "AgentscopeRuntimeMemoryService", - "PersonalMemoryService", - "TaskMemoryService", -] diff --git a/reme_ai/service/agentscope_runtime_memory_service.py b/reme_ai/service/agentscope_runtime_memory_service.py deleted file mode 100644 index 76b50c61..00000000 --- a/reme_ai/service/agentscope_runtime_memory_service.py +++ /dev/null @@ -1,160 +0,0 @@ -"""Base memory service for Agentscope runtime integration. - -This module provides the abstract base class AgentscopeRuntimeMemoryService -which defines the interface for memory services that integrate with -Agentscope runtime. Concrete implementations should inherit from this class -and implement the abstract methods. -""" - -from abc import abstractmethod, ABC -from typing import Optional, Dict, Any - -from pydantic import Field - -from reme_ai.main import ReMeApp - - -class AgentscopeRuntimeMemoryService(ABC): - """Abstract base class for memory services integrated with Agentscope runtime. - - This class provides a common interface for memory services and manages - the underlying ReMeApp instance and session-to-memory-id mappings. - Subclasses must implement the abstract methods to provide specific - memory management functionality. - """ - - def __init__(self): - """Initialize the memory service. - - Creates a new ReMeApp instance and initializes the session-to-memory-id - mapping dictionary. - """ - self.app = ReMeApp() - self.session_id_dict: dict = {} - - def add_session_memory_id(self, session_id: str, memory_id): - """Add a memory ID to a session's memory list. - - Associates a memory_id with a session_id by adding it to the - session's memory list. If the session doesn't exist, it will - be created. - - Args: - session_id: The session identifier. - memory_id: The memory identifier to associate with the session. - """ - if session_id not in self.session_id_dict: - self.session_id_dict[session_id] = [] - - self.session_id_dict[session_id].append(memory_id) - - @abstractmethod - async def start(self) -> None: - """Starts the service, initializing any necessary resources or - connections.""" - - @abstractmethod - async def stop(self) -> None: - """Stops the service, releasing any acquired resources.""" - - @abstractmethod - async def health(self) -> bool: - """ - Checks the health of the service. - - Returns: - True if the service is healthy, False otherwise. - """ - - async def __aenter__(self): - """Async context manager entry. - - Starts the service when entering an async context. - - Returns: - The service instance. - """ - await self.start() - return self - - async def __aexit__(self, exc_type, exc_val, exc_tb): - """Async context manager exit. - - Stops the service when exiting an async context. - - Args: - exc_type: Exception type if an exception occurred. - exc_val: Exception value if an exception occurred. - exc_tb: Exception traceback if an exception occurred. - - Returns: - False to propagate exceptions, True to suppress them. - """ - await self.stop() - return False - - @abstractmethod - async def add_memory( - self, - user_id: str, - messages: list, - session_id: Optional[str] = None, - ) -> None: - """ - Adds messages to the memory service. - - Args: - user_id: The user id. - messages: The messages to add. - session_id: The session id, which is optional. - """ - - @abstractmethod - async def search_memory( - self, - user_id: str, - messages: list, - filters: Optional[Dict[str, Any]] = Field( - description="Associated filters for the messages, " "such as top_k, score etc.", - default=None, - ), - ) -> list: - """ - Searches messages from the memory service. - - Args: - user_id: The user id. - messages: The user query or the query with history messages, - both in the format of list of messages. If messages is a list, - the search will be based on the content of the last message. - filters: The filters used to search memory - """ - - @abstractmethod - async def list_memory( - self, - user_id: str, - filters: Optional[Dict[str, Any]] = Field( - description="Associated filters for the messages, " "such as top_k, score etc.", - default=None, - ), - ) -> list: - """ - Lists the memory items for a given user with filters, such as - page_num, page_size, etc. - - Args: - user_id: The user id. - filters: The filters for the memory items. - """ - - @abstractmethod - async def delete_memory( - self, - user_id: str, - session_id: Optional[str] = None, - ) -> None: - """ - Deletes the memory items for a given user with certain session id, - or all the memory items for a given user. - """ diff --git a/reme_ai/service/personal_memory_service.py b/reme_ai/service/personal_memory_service.py deleted file mode 100644 index c97dd465..00000000 --- a/reme_ai/service/personal_memory_service.py +++ /dev/null @@ -1,216 +0,0 @@ -"""Personal memory service for managing personalized memories. - -This module provides the PersonalMemoryService class which extends the base -AgentscopeRuntimeMemoryService to handle personal memory operations. -It supports creating, retrieving, listing, and deleting personal memories -using flow-based execution. -""" - -import asyncio -from typing import Optional, Dict, Any, List - -from flowllm.core.schema import FlowResponse -from loguru import logger -from pydantic import Field, BaseModel - -from reme_ai.schema.memory import PersonalMemory -from reme_ai.service.agentscope_runtime_memory_service import AgentscopeRuntimeMemoryService - - -class PersonalMemoryService(AgentscopeRuntimeMemoryService): - """Service for managing personalized memories. - - PersonalMemoryService empowers you to generate, retrieve, and share - customized memories. Leveraging advanced LLM, embedding, and vector store - technologies, it builds a comprehensive memory system with intelligent, - context- and time-aware memory management. - """ - - async def start(self): - """Start the personal memory service. - - Returns: - The result of starting the underlying application. - """ - return await self.app.async_start() - - async def stop(self) -> None: - """Stop the personal memory service. - - Releases resources and stops the underlying application. - """ - return await self.app.async_stop() - - async def health(self) -> bool: - """Check the health status of the service. - - Returns: - True if the service is healthy, False otherwise. - """ - return True - - async def add_memory(self, user_id: str, messages: list, session_id: Optional[str] = None) -> None: - """Add personal memory from messages. - - Processes the provided messages and creates personal memories using - the summary_personal_memory flow. The created memories are associated - with the given session_id. - - Args: - user_id: The user identifier. - messages: List of messages (dict or BaseModel instances) to process. - session_id: Optional session identifier to associate with the memory. - """ - new_messages: List[dict] = [] - for message in messages: - if isinstance(message, dict): - new_messages.append(message) - elif isinstance(message, BaseModel): - new_messages.append(message.model_dump()) - else: - raise ValueError(f"Invalid message type={type(message)}") - - kwargs = { - "workspace_id": user_id, - "trajectories": [ - {"messages": new_messages, "score": 1.0}, - ], - } - - result: FlowResponse = await self.app.async_execute_flow(name="summary_personal_memory", **kwargs) - memory_list: List[PersonalMemory] = result.metadata.get("memory_list", []) - for memory in memory_list: - memory_id = memory.memory_id - self.add_session_memory_id(session_id, memory_id) - logger.info(f"[personal_memory_service] user_id={user_id} session_id={session_id} add memory: {memory}") - - async def search_memory( - self, - user_id: str, - messages: list, - filters: Optional[Dict[str, Any]] = Field( - description="Associated filters for the messages, " "such as top_k, score etc.", - default=None, - ), - ) -> list: - """Search for personal memories matching the given messages. - - Searches the memory store for personal memories relevant to the provided - messages using the retrieve_personal_memory flow. The query is extracted - from the last message in the messages list. - - Args: - user_id: The user identifier. - messages: List of messages (dict or BaseModel instances) to search with. - filters: Optional filters including top_k for controlling search results. - - Returns: - List containing the search result answer. - """ - new_messages: List[dict] = [] - for message in messages: - if isinstance(message, dict): - new_messages.append(message) - elif isinstance(message, BaseModel): - new_messages.append(message.model_dump()) - else: - raise ValueError(f"Invalid message type={type(message)}") - - # Extract query from the last message - query = new_messages[-1]["content"] if messages else "" - - kwargs = { - "workspace_id": user_id, - "query": query, - "top_k": filters.get("top_k", 1) if filters else 1, - } - - result: FlowResponse = await self.app.async_execute_flow(name="retrieve_personal_memory", **kwargs) - logger.info(f"[personal_memory_service] user_id={user_id} search result: {result.model_dump_json()}") - - return [result.answer] - - async def list_memory( - self, - user_id: str, - filters: Optional[Dict[str, Any]] = Field( - description="Associated filters for the messages, " "such as top_k, score etc.", - default=None, - ), - ) -> list: - """List all personal memories for a user. - - Retrieves all personal memories associated with the given user_id - from the vector store. - - Args: - user_id: The user identifier. - filters: Optional filters (currently not used but kept for API consistency). - - Returns: - List of memory items for the user. - """ - result = await self.app.async_execute_flow(name="vector_store", workspace_id=user_id, action="list") - logger.info(f"[personal_memory_service] list_memory result: {result}") - - result = result.metadata["action_result"] - for i, line in enumerate(result): - logger.info(f"[personal_memory_service] list memory.{i}={line}") - return result - - async def delete_memory(self, user_id: str, session_id: Optional[str] = None) -> None: - """Delete personal memories for a user session. - - Deletes all memories associated with the given session_id for the user. - If no session_id is provided or no memories exist for the session, - no deletion is performed. - - Args: - user_id: The user identifier. - session_id: Optional session identifier. If provided, only memories - associated with this session will be deleted. - """ - delete_ids = self.session_id_dict.get(session_id, []) - if not delete_ids: - return - - result = await self.app.async_execute_flow( - name="vector_store", - workspace_id=user_id, - action="delete_ids", - memory_ids=delete_ids, - ) - result = result.metadata["action_result"] - logger.info(f"[personal_memory_service] delete memory result={result}") - - -async def main(): - """Main function for testing the PersonalMemoryService. - - Demonstrates the usage of PersonalMemoryService by adding, searching, - listing, and deleting personal memories. - """ - async with PersonalMemoryService() as service: - logger.info("========== start personal memory service ==========") - - await service.add_memory( - user_id="u_12345", - messages=[{"content": "I really enjoy playing tennis on weekends"}], - session_id="s_123456", - ) - - await service.search_memory( - user_id="u_12345", - messages=[{"content": "What do I like to do for fun?"}], - filters={"top_k": 1}, - ) - - await service.list_memory(user_id="u_12345") - await service.delete_memory(user_id="u_12345", session_id="s_123456") - await service.list_memory(user_id="u_12345") - - logger.info("========== end personal memory service ==========") - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/reme_ai/service/task_memory_service.py b/reme_ai/service/task_memory_service.py deleted file mode 100644 index a8088cb8..00000000 --- a/reme_ai/service/task_memory_service.py +++ /dev/null @@ -1,212 +0,0 @@ -"""Task memory service for managing task-oriented memories. - -This module provides the TaskMemoryService class which extends the base -AgentscopeRuntimeMemoryService to handle task-related memory operations. -It supports creating, retrieving, listing, and deleting task memories -using flow-based execution. -""" - -import asyncio -from typing import Optional, Dict, Any, List - -from flowllm.core.schema import FlowResponse -from loguru import logger -from pydantic import Field, BaseModel - -from reme_ai.schema.memory import TaskMemory -from reme_ai.service.agentscope_runtime_memory_service import AgentscopeRuntimeMemoryService - - -class TaskMemoryService(AgentscopeRuntimeMemoryService): - """Service for managing task-oriented memories. - - TaskMemoryService helps efficiently manage and schedule task-related memories, - enhancing both the accuracy and efficiency of task execution. Powered by LLM - capabilities, it supports flexible creation, retrieval, update, and deletion - of memories across diverse task scenarios. - """ - - async def start(self): - """Start the task memory service. - - Returns: - The result of starting the underlying application. - """ - return await self.app.async_start() - - async def stop(self) -> None: - """Stop the task memory service. - - Releases resources and stops the underlying application. - """ - return await self.app.async_stop() - - async def health(self) -> bool: - """Check the health status of the service. - - Returns: - True if the service is healthy, False otherwise. - """ - return True - - async def add_memory(self, user_id: str, messages: list, session_id: Optional[str] = None) -> None: - """Add task memory from messages. - - Processes the provided messages and creates task memories using - the summary_task_memory flow. The created memories are associated - with the given session_id. - - Args: - user_id: The user identifier. - messages: List of messages (dict or BaseModel instances) to process. - session_id: Optional session identifier to associate with the memory. - """ - new_messages: List[dict] = [] - for message in messages: - if isinstance(message, dict): - new_messages.append(message) - elif isinstance(message, BaseModel): - new_messages.append(message.model_dump()) - else: - raise ValueError(f"Invalid message type={type(message)}") - - kwargs = { - "workspace_id": user_id, - "trajectories": [ - {"messages": new_messages, "score": 1.0}, - ], - } - - result: FlowResponse = await self.app.async_execute_flow(name="summary_task_memory", **kwargs) - memory_list: List[TaskMemory] = result.metadata.get("memory_list", []) - for memory in memory_list: - memory_id = memory.memory_id - self.add_session_memory_id(session_id, memory_id) - logger.info(f"[task_memory_service] user_id={user_id} session_id={session_id} add memory: {memory}") - - async def search_memory( - self, - user_id: str, - messages: list, - filters: Optional[Dict[str, Any]] = Field( - description="Associated filters for the messages, " "such as top_k, score etc.", - default=None, - ), - ) -> list: - """Search for task memories matching the given messages. - - Searches the memory store for task memories relevant to the provided - messages using the retrieve_task_memory flow. - - Args: - user_id: The user identifier. - messages: List of messages (dict or BaseModel instances) to search with. - filters: Optional filters including top_k for controlling search results. - - Returns: - List containing the search result answer. - """ - new_messages: List[dict] = [] - for message in messages: - if isinstance(message, dict): - new_messages.append(message) - elif isinstance(message, BaseModel): - new_messages.append(message.model_dump()) - else: - raise ValueError(f"Invalid message type={type(message)}") - - kwargs = { - "workspace_id": user_id, - "messages": new_messages, - "top_k": filters.get("top_k", 1) if filters else 1, - } - - result: FlowResponse = await self.app.async_execute_flow(name="retrieve_task_memory", **kwargs) - logger.info(f"[task_memory_service] user_id={user_id} add result: {result.model_dump_json()}") - - return [result.answer] - - async def list_memory( - self, - user_id: str, - filters: Optional[Dict[str, Any]] = Field( - description="Associated filters for the messages, " "such as top_k, score etc.", - default=None, - ), - ) -> list: - """List all task memories for a user. - - Retrieves all task memories associated with the given user_id - from the vector store. - - Args: - user_id: The user identifier. - filters: Optional filters (currently not used but kept for API consistency). - - Returns: - List of memory items for the user. - """ - result = await self.app.async_execute_flow(name="vector_store", workspace_id=user_id, action="list") - print("list_memory result:", result) - - result = result.metadata["action_result"] - for i, line in enumerate(result): - logger.info(f"[task_memory_service] list memory.{i}={line}") - return result - - async def delete_memory(self, user_id: str, session_id: Optional[str] = None) -> None: - """Delete task memories for a user session. - - Deletes all memories associated with the given session_id for the user. - If no session_id is provided or no memories exist for the session, - no deletion is performed. - - Args: - user_id: The user identifier. - session_id: Optional session identifier. If provided, only memories - associated with this session will be deleted. - """ - delete_ids = self.session_id_dict.get(session_id, []) - if not delete_ids: - return - - result = await self.app.async_execute_flow( - name="vector_store", - workspace_id=user_id, - action="delete_ids", - memory_ids=delete_ids, - ) - result = result.metadata["action_result"] - logger.info(f"[task_memory_service] delete memory result={result}") - - -async def main(): - """Main function for testing the TaskMemoryService. - - Demonstrates the usage of TaskMemoryService by adding, searching, - listing, and deleting task memories. - """ - async with TaskMemoryService() as service: - logger.info("========== start task memory service ==========") - - await service.add_memory( - user_id="u_123456", - messages=[{"content": "please use web search tool to search financial news:"}], - session_id="s_123456", - ) - - await service.search_memory( - user_id="u_123456", - messages=[{"content": "please use web search tool to search financial news"}], - filters={"top_k": 1}, - ) - - await service.list_memory(user_id="u_123456") - await service.delete_memory(user_id="u_123456", session_id="s_123456") - await service.list_memory(user_id="u_123456") - - logger.info("========== end task memory service ==========") - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/reme_ai/summary/__init__.py b/reme_ai/summary/__init__.py deleted file mode 100644 index 47688a15..00000000 --- a/reme_ai/summary/__init__.py +++ /dev/null @@ -1,20 +0,0 @@ -"""Summary operations' module. - -This module provides summary operations for different types of memories: -- Personal memory summary operations -- Task memory summary operations -- Tool memory summary operations -- Working memory summary operations -""" - -from . import personal -from . import task -from . import tool -from . import working - -__all__ = [ - "personal", - "task", - "tool", - "working", -] diff --git a/reme_ai/summary/personal/__init__.py b/reme_ai/summary/personal/__init__.py deleted file mode 100644 index 2c5533b5..00000000 --- a/reme_ai/summary/personal/__init__.py +++ /dev/null @@ -1,30 +0,0 @@ -"""Personal memory summary operations module. - -This module provides operations for processing personal memories, including: -- Filtering messages based on information content -- Extracting observations from chat messages -- Generating reflection subjects -- Updating insights based on new observations -- Detecting and handling contradictions and redundancies -- Loading today's memories for deduplication -""" - -from reme_ai.summary.personal.contra_repeat_op import ContraRepeatOp -from reme_ai.summary.personal.get_observation_op import GetObservationOp -from reme_ai.summary.personal.get_observation_with_time_op import GetObservationWithTimeOp -from reme_ai.summary.personal.get_reflection_subject_op import GetReflectionSubjectOp -from reme_ai.summary.personal.info_filter_op import InfoFilterOp -from reme_ai.summary.personal.load_today_memory_op import LoadTodayMemoryOp -from reme_ai.summary.personal.long_contra_repeat_op import LongContraRepeatOp -from reme_ai.summary.personal.update_insight_op import UpdateInsightOp - -__all__ = [ - "ContraRepeatOp", - "GetObservationOp", - "GetObservationWithTimeOp", - "GetReflectionSubjectOp", - "InfoFilterOp", - "LoadTodayMemoryOp", - "LongContraRepeatOp", - "UpdateInsightOp", -] diff --git a/reme_ai/summary/personal/contra_repeat_op.py b/reme_ai/summary/personal/contra_repeat_op.py deleted file mode 100644 index aec839b2..00000000 --- a/reme_ai/summary/personal/contra_repeat_op.py +++ /dev/null @@ -1,160 +0,0 @@ -"""Module for detecting and handling contradictory and repetitive memories. - -This module provides the ContraRepeatOp class which processes memory nodes -to identify and handle contradictory and repetitive information. It collects -observation memories from context, constructs prompts for language model -analysis, parses responses to detect contradictions or redundancies, and -filters the processed memories accordingly. -""" - -import json -import re -from typing import List, Tuple - -from flowllm.core.context import C -from flowllm.core.enumeration import Role -from flowllm.core.op import BaseAsyncOp -from flowllm.core.schema import Message -from loguru import logger - -from reme_ai.schema.memory import BaseMemory - - -@C.register_op() -class ContraRepeatOp(BaseAsyncOp): - """ - The `ContraRepeatOp` class specializes in processing memory nodes to identify and handle - contradictory and repetitive information. It extends the base functionality of `BaseAsyncOp`. - - Responsibilities: - - Collects observation memories from context. - - Constructs a prompt with these observations for language model analysis. - - Parses the model's response to detect contradictions or redundancies. - - Filters and returns the processed memories. - """ - - file_path: str = __file__ - - async def async_execute(self): - """ - Executes the primary routine of the ContraRepeatOp which involves: - 1. Gets memory list from context - 2. Constructs a prompt with these memories for language model analysis - 3. Parses the model's response to detect contradictions or redundancies - 4. Filters and returns the processed memories - """ - # Get memory list from context - standardized key - memory_list: List[BaseMemory] = [] - memory_list.extend(self.context.get("observation_memories", [])) - memory_list.extend(self.context.get("observation_memories_with_time", [])) - memory_list.extend(self.context.get("today_memories", [])) - - self.context.response.metadata["memory_list"] = memory_list - - if not memory_list: - logger.info("memory_list is empty!") - self.context.response.metadata["deleted_memory_ids"] = [] - return - - # Get operation parameters - contra_repeat_max_count: int = self.op_params.get("contra_repeat_max_count", 50) - enable_contra_repeat: bool = self.op_params.get("enable_contra_repeat", True) - - if not enable_contra_repeat: - logger.warning("contra_repeat is not enabled!") - self.context.response.metadata["deleted_memory_ids"] = [] - return - - # Sort and limit memories by count - sorted_memories = sorted(memory_list, key=lambda x: x.time_created, reverse=True)[:contra_repeat_max_count] - - if len(sorted_memories) <= 1: - logger.info("sorted_memories.size<=1, stop.") - self.context.response.metadata["memory_list"] = sorted_memories - self.context.response.metadata["deleted_memory_ids"] = [] - return - - # Build prompt - user_query_list = [] - for i, memory in enumerate(sorted_memories): - user_query_list.append(f"{i + 1} {memory.content}") - - user_name = self.context.get("user_name", "user") - - # Create prompt using the new pattern - system_prompt = self.prompt_format( - prompt_name="contra_repeat_system", - num_obs=len(user_query_list), - user_name=user_name, - ) - few_shot = self.prompt_format(prompt_name="contra_repeat_few_shot", user_name=user_name) - user_query = self.prompt_format( - prompt_name="contra_repeat_user_query", - user_query="\n".join(user_query_list), - ) - - full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}" - logger.info(f"contra_repeat_prompt={full_prompt}") - - # Call LLM - response = await self.llm.achat([Message(role=Role.USER, content=full_prompt)]) - - # Return if empty - if not response or not response.content: - logger.warning("Empty response from LLM") - self.context.response.metadata["memory_list"] = sorted_memories - self.context.response.metadata["deleted_memory_ids"] = [] - return - - response_text = response.content - logger.info(f"contra_repeat_response={response_text}") - - # Parse response and filter memories - filtered_memories, deleted_memory_ids = self._parse_and_filter_memories(response_text, sorted_memories) - - # Update context with filtered memories and deleted memory IDs - standardized keys - self.context.response.metadata["memory_list"] = filtered_memories - self.context.response.metadata["deleted_memory_ids"] = deleted_memory_ids - logger.info(f"Filtered {len(memory_list)} memories to {len(filtered_memories)} memories") - logger.info(f"Deleted memory IDs: {json.dumps(deleted_memory_ids, indent=2)}") - - @staticmethod - def _parse_and_filter_memories(response_text: str, memories: List[BaseMemory]) -> Tuple[ - List[BaseMemory], - List[str], - ]: - """Parse LLM response and filter memories based on contradiction/containment analysis""" - - # Parse the response to extract judgments - pattern = r"<(\d+)>\s*<(矛盾|被包含|无|Contradiction|Contained|None)>" - matches = re.findall(pattern, response_text, re.IGNORECASE) - - if not matches: - logger.warning("No valid judgments found in response") - return memories, [] - - # Create a set of indices to remove (contradictory or contained memories) - indices_to_remove = set() - deleted_memory_ids = [] - - for idx_str, judgment in matches: - try: - idx = int(idx_str) - 1 # Convert to 0-based index - if idx >= len(memories): - logger.warning(f"Invalid index {idx} for memories list of length {len(memories)}") - continue - - judgment_lower = judgment.lower() - if judgment_lower in ["矛盾", "contradiction", "被包含", "contained"]: - indices_to_remove.add(idx) - deleted_memory_ids.append(memories[idx].memory_id) - logger.info(f"Marking memory {idx + 1} for removal: {judgment} - {memories[idx].content[:100]}...") - - except ValueError: - logger.warning(f"Invalid index format: {idx_str}") - continue - - # Filter out the memories marked for removal - filtered_memories = [memory for i, memory in enumerate(memories) if i not in indices_to_remove] - - return filtered_memories, deleted_memory_ids diff --git a/reme_ai/summary/personal/get_observation_op.py b/reme_ai/summary/personal/get_observation_op.py deleted file mode 100644 index e5d7ca85..00000000 --- a/reme_ai/summary/personal/get_observation_op.py +++ /dev/null @@ -1,163 +0,0 @@ -"""Module for generating observations from chat messages. - -This module provides the GetObservationOp class which extracts personal -observations from chat messages. It filters messages to exclude those with -time-related keywords and uses LLM-based extraction to generate structured -observation memories from the filtered messages. -""" - -import re -from typing import List - -from flowllm.core.context import C -from flowllm.core.op import BaseAsyncOp -from flowllm.core.schema import Message -from loguru import logger - -from reme_ai.schema.memory import BaseMemory, PersonalMemory -from reme_ai.utils.datetime_handler import DatetimeHandler - - -@C.register_op() -class GetObservationOp(BaseAsyncOp): - """ - A specialized operation class to generate observations from chat messages using BaseAsyncOp. - """ - - file_path: str = __file__ - - async def async_execute(self): - """Extract personal observations from chat messages""" - # Get messages from context - guaranteed to exist by flow input - messages: List[Message] = self.context.messages - if not messages: - logger.warning("No messages found in context") - return - - # Filter messages - exclude those with time-related keywords - filtered_messages = self._filter_messages(messages) - if not filtered_messages: - logger.warning("No messages left after filtering") - self.context.observation_memories = [] - return - - logger.info(f"Extracting observations from {len(filtered_messages)} filtered messages") - - # Extract observations using LLM - observation_memories = await self._extract_observations_from_messages(filtered_messages) - - # Store results in context using standardized key - self.context.observation_memories = observation_memories - logger.info(f"Generated {len(observation_memories)} observation memories") - - def _filter_messages(self, messages: List[Message]) -> List[Message]: - """ - Filters the chat messages to exclude those containing time-related keywords. - - Args: - messages: List of messages to filter - - Returns: - List[Message]: A list of filtered messages without time keywords. - """ - filtered_messages = [] - for msg in messages: - if not DatetimeHandler.has_time_word(query=msg.content, language=self.language): - filtered_messages.append(msg) - - logger.info(f"Filtered messages from {len(messages)} to {len(filtered_messages)}") - return filtered_messages - - async def _extract_observations_from_messages(self, filtered_messages: List[Message]) -> List[BaseMemory]: - """Extract observations from filtered messages using LLM""" - user_name = self.context.get("user_name", "user") - - # Build prompt for observation extraction - user_query_list = [] - for i, msg in enumerate(filtered_messages): - user_query_list.append(f"{i + 1} {user_name}: {msg.content}") - - # Create prompt using the prompt format method - system_prompt = self.prompt_format( - prompt_name="get_observation_system", - num_obs=len(user_query_list), - user_name=user_name, - ) - few_shot = self.prompt_format(prompt_name="get_observation_few_shot", user_name=user_name) - user_query = self.prompt_format( - prompt_name="get_observation_user_query", - user_query="\n".join(user_query_list), - user_name=user_name, - ) - - full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}" - logger.info(f"get_observation_prompt={full_prompt}") - - def parse_observations(message: Message) -> List[BaseMemory]: - """Parse LLM response and create observation memories""" - response_text = message.content - logger.info(f"get_observation_response={response_text}") - - # Parse observations using class method - parsed_observations = GetObservationOp.parse_observation_response(response_text) - - observation_memories = [] - for obs in parsed_observations: - idx = obs["index"] - 1 # Convert to 0-based index - if idx >= len(filtered_messages): - logger.warning(f"Invalid index {idx} for messages list of length {len(filtered_messages)}") - continue - - # Create observation memory - observation = PersonalMemory( - workspace_id=self.context.workspace_id, - when_to_use=obs["keywords"], - content=obs["content"], - target=user_name, - author=self.llm.model_name, - metadata={ - "keywords": obs["keywords"], - "source_message": filtered_messages[idx].content, - "observation_type": "personal_info", - }, - ) - observation_memories.append(observation) - logger.info(f"Created observation: {obs['content'][:50]}...") - - return observation_memories - - # Use LLM chat with callback function - return await self.llm.achat(messages=[Message(content=full_prompt)], callback_fn=parse_observations) - - @staticmethod - def parse_observation_response(response_text: str) -> List[dict]: - """Parse observation response to extract structured data""" - # Pattern to match both Chinese and English observation formats - pattern = r"信息:<(\d+)>\s*<>\s*<([^<>]+)>\s*<([^<>]*)>|Information:\s*<(\d+)>\s*<>\s*<([^<>]+)>\s*<([^<>]*)>" - matches = re.findall(pattern, response_text, re.IGNORECASE | re.MULTILINE) - - observations = [] - for match in matches: - # Handle both Chinese and English patterns - if match[0]: # Chinese pattern - idx_str, content, keywords = match[0], match[1], match[2] - else: # English pattern - idx_str, content, keywords = match[3], match[4], match[5] - - try: - idx = int(idx_str) - # Skip if content indicates no meaningful observation - content_lower = content.lower().strip() - if content_lower not in ["无", "none", "", "repeat"]: - observations.append( - { - "index": idx, - "content": content.strip(), - "keywords": keywords.strip() if keywords else "", - }, - ) - except ValueError: - logger.warning(f"Invalid index format: {idx_str}") - continue - - return observations diff --git a/reme_ai/summary/personal/get_observation_with_time_op.py b/reme_ai/summary/personal/get_observation_with_time_op.py deleted file mode 100644 index e7c2b91a..00000000 --- a/reme_ai/summary/personal/get_observation_with_time_op.py +++ /dev/null @@ -1,186 +0,0 @@ -"""Module for extracting observations with time information from chat messages. - -This module provides the GetObservationWithTimeOp class which extracts -personal observations with time information from chat messages. It filters -messages to only include those with time-related keywords and uses LLM-based -extraction to generate structured observation memories with time information -from the filtered messages. -""" - -import re -from typing import List - -from flowllm.core.context import C -from flowllm.core.op import BaseAsyncOp -from flowllm.core.schema import Message -from loguru import logger - -from reme_ai.schema.memory import BaseMemory, PersonalMemory -from reme_ai.utils.datetime_handler import DatetimeHandler - - -@C.register_op() -class GetObservationWithTimeOp(BaseAsyncOp): - """ - A specialized operation class to extract observations with time information from chat messages using BaseAsyncOp. - """ - - file_path: str = __file__ - - async def async_execute(self): - """Extract personal observations with time information from chat messages""" - # Get messages from context - guaranteed to exist by flow input - messages: List[Message] = self.context.messages - if not messages: - logger.warning("No messages found in context") - return - - # Filter messages - only include those with time-related keywords - filtered_messages = self._filter_messages(messages) - if not filtered_messages: - logger.warning("No messages with time keywords found") - self.context.observation_memories_with_time = [] - return - - logger.info(f"Extracting observations with time from {len(filtered_messages)} filtered messages") - - # Extract observations using LLM - observation_memories_with_time = await self._extract_observations_with_time_from_messages(filtered_messages) - - # Store results in context using standardized key - self.context.observation_memories_with_time = observation_memories_with_time - logger.info(f"Generated {len(observation_memories_with_time)} observation memories with time") - - def _filter_messages(self, messages: List[Message]) -> List[Message]: - """ - Filters the chat messages to only include those containing time-related keywords. - - Args: - messages: List of messages to filter - - Returns: - List[Message]: A list of filtered messages that mention time. - """ - filtered_messages = [] - for msg in messages: - if DatetimeHandler.has_time_word(query=msg.content, language=self.language): - filtered_messages.append(msg) - - logger.info(f"Filtered messages from {len(messages)} to {len(filtered_messages)}") - return filtered_messages - - async def _extract_observations_with_time_from_messages(self, filtered_messages: List[Message]) -> List[BaseMemory]: - """Extract observations with time information from filtered messages using LLM""" - user_name = self.context.get("user_name", "user") - - # Build prompt for observation extraction with time - user_query_list = [] - for i, msg in enumerate(filtered_messages): - # Create a DatetimeHandler instance for each message's timestamp and format it - dt_handler = DatetimeHandler(dt=msg.time_created) - - # Get time format from prompt configuration - time_format = self.prompt_format(prompt_name="time_string_format") - dt = dt_handler.string_format(string_format=time_format, language=self.language) - - # Append formatted timestamp-query pairs to the user_query_list - colon = self._get_colon_word() - user_query_list.append(f"{i + 1} {dt} {user_name}{colon}{msg.content}") - - # Create prompt using the prompt format method - system_prompt = self.prompt_format( - prompt_name="get_observation_with_time_system", - num_obs=len(user_query_list), - user_name=user_name, - ) - few_shot = self.prompt_format(prompt_name="get_observation_with_time_few_shot", user_name=user_name) - user_query = self.prompt_format( - prompt_name="get_observation_with_time_user_query", - user_query="\n".join(user_query_list), - user_name=user_name, - ) - - full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}" - logger.info(f"get_observation_with_time_prompt={full_prompt}") - - def parse_observations(message: Message) -> List[BaseMemory]: - """Parse LLM response and create observation memories with time""" - response_text = message.content - logger.info(f"get_observation_with_time_response={response_text}") - - # Parse observations using class method - parsed_observations = GetObservationWithTimeOp.parse_observation_with_time_response(response_text) - - observation_memories = [] - for obs in parsed_observations: - idx = obs["index"] - 1 # Convert to 0-based index - if idx >= len(filtered_messages): - logger.warning(f"Invalid index {idx} for messages list of length {len(filtered_messages)}") - continue - - # Create observation memory - observation = PersonalMemory( - workspace_id=self.context.workspace_id, - when_to_use=obs["keywords"], - content=obs["content"], - target=user_name, - author=getattr(self.llm, "model_name", "system"), - metadata={ - "keywords": obs["keywords"], - "time_info": obs.get("time_info", ""), - "source_message": filtered_messages[idx].content, - "observation_type": "personal_info_with_time", - }, - ) - observation_memories.append(observation) - logger.info(f"Created observation with time: {obs['content'][:50]}...") - - return observation_memories - - # Use LLM chat with callback function - return await self.llm.achat(messages=[Message(content=full_prompt)], callback_fn=parse_observations) - - def _get_colon_word(self) -> str: - """Get language-specific colon word""" - colon_dict = {"zh": ":", "cn": ":", "en": ": "} - return colon_dict.get(self.language, ": ") - - @staticmethod - def parse_observation_with_time_response(response_text: str) -> List[dict]: - """Parse observation with time response to extract structured data""" - # Pattern to match both Chinese and English observation formats with time information - # Chinese: 信息:<1> <时间信息或不输出> <明确的重要信息或"无"> <关键词> - # English: Information: <1>