diff --git a/backend/app/api/graph.py b/backend/app/api/graph.py index cb7f0dff..da3da140 100644 --- a/backend/app/api/graph.py +++ b/backend/app/api/graph.py @@ -14,6 +14,7 @@ from ..services.ontology_generator import OntologyGenerator from ..services.graph_builder import GraphBuilderService from ..services.text_processor import TextProcessor from ..utils.file_parser import FileParser +from ..utils.llm_client import LLMClient from ..utils.logger import get_logger from ..models.task import TaskManager, TaskStatus from ..models.project import ProjectManager, ProjectStatus @@ -30,8 +31,6 @@ def allowed_file(filename: str) -> bool: return ext in Config.ALLOWED_EXTENSIONS -# ============== 项目管理接口 ============== - @graph_bp.route('/project/', methods=['GET']) def get_project(project_id: str): """ @@ -213,7 +212,15 @@ def generate_ontology(): # Generate ontology definition using LLM logger.info("Calling LLM to generate ontology definition...") - generator = OntologyGenerator() + generator = OntologyGenerator( + llm_client=LLMClient( + component="ontology_generator", + metadata={ + "project_id": project.project_id, + "phase": "ontology_generation", + }, + ) + ) ontology = generator.generate( document_texts=document_texts, simulation_requirement=simulation_requirement, diff --git a/backend/app/api/report.py b/backend/app/api/report.py index 1111d281..6715f5aa 100644 --- a/backend/app/api/report.py +++ b/backend/app/api/report.py @@ -134,6 +134,7 @@ def generate_report(): agent = ReportAgent( graph_id=graph_id, simulation_id=simulation_id, + project_id=state.project_id, simulation_requirement=simulation_requirement ) @@ -540,6 +541,7 @@ def chat_with_report_agent(): agent = ReportAgent( graph_id=graph_id, simulation_id=simulation_id, + project_id=state.project_id, simulation_requirement=simulation_requirement ) diff --git a/backend/app/api/simulation.py b/backend/app/api/simulation.py index 48b6e99f..b999a825 100644 --- a/backend/app/api/simulation.py +++ b/backend/app/api/simulation.py @@ -1397,6 +1397,7 @@ def generate_profiles(): entity_types = data.get('entity_types') use_llm = data.get('use_llm', True) platform = data.get('platform', 'reddit') + project_id = data.get('project_id') reader = ZepEntityReader() filtered = reader.filter_defined_entities( @@ -1414,7 +1415,8 @@ def generate_profiles(): generator = OasisProfileGenerator() profiles = generator.generate_profiles_from_entities( entities=filtered.entities, - use_llm=use_llm + use_llm=use_llm, + project_id=project_id, ) if platform == "reddit": diff --git a/backend/app/services/oasis_profile_generator.py b/backend/app/services/oasis_profile_generator.py index 4d2c5a97..c7595543 100644 --- a/backend/app/services/oasis_profile_generator.py +++ b/backend/app/services/oasis_profile_generator.py @@ -20,6 +20,7 @@ from zep_cloud.client import Zep from ..config import Config from ..utils.logger import get_logger +from ..utils.llm_cost import create_tracked_chat_completion from .zep_entity_reader import EntityNode, ZepEntityReader logger = get_logger('mirofish.oasis_profile') @@ -196,6 +197,10 @@ class OasisProfileGenerator: api_key=self.api_key, base_url=self.base_url ) + self._runtime_metadata: Dict[str, Any] = { + "component": "oasis_profile_generator", + "phase": "generate_profiles", + } # Kết nối lên trên ZEP Database Search Context self.zep_api_key = zep_api_key or Config.ZEP_API_KEY @@ -525,15 +530,16 @@ class OasisProfileGenerator: for attempt in range(max_attempts): try: - response = self.client.chat.completions.create( + response = create_tracked_chat_completion( + client=self.client, model=self.model_name, messages=[ {"role": "system", "content": self._get_system_prompt(is_individual)}, {"role": "user", "content": prompt} ], response_format={"type": "json_object"}, - temperature=0.7 - (attempt * 0.1) # Giảm tính sáng tạo ngẫu nhiên đi một chút mỗi khi fail để tăng khả năng thành công ở vòng tiếp theo - # Thả Rông Max tokens + temperature=0.7 - (attempt * 0.1), # Giảm tính sáng tạo ngẫu nhiên đi một chút mỗi khi fail để tăng khả năng thành công ở vòng tiếp theo + metadata=self._runtime_metadata, ) content = response.choices[0].message.content @@ -854,7 +860,9 @@ Quan trọng: graph_id: Optional[str] = None, parallel_count: int = 5, realtime_output_path: Optional[str] = None, - output_platform: str = "reddit" + output_platform: str = "reddit", + simulation_id: Optional[str] = None, + project_id: Optional[str] = None, ) -> List[OasisAgentProfile]: """ Khởi tạo hàng loạt các Agent Profile từ các thực thể (Hỗ trợ Gen đa luồng song song) @@ -877,6 +885,14 @@ Quan trọng: # Lưu Graph ID lại cho Zep xử lý search if graph_id: self.graph_id = graph_id + + self._runtime_metadata = { + "component": "oasis_profile_generator", + "phase": "generate_profiles", + "simulation_id": simulation_id, + "project_id": project_id, + "platform": output_platform, + } total = len(entities) profiles = [None] * total # Cấp trước 1 mảng để giữ đúng thứ tự Index diff --git a/backend/app/services/ontology_generator.py b/backend/app/services/ontology_generator.py index 3303ad1f..cc44a899 100644 --- a/backend/app/services/ontology_generator.py +++ b/backend/app/services/ontology_generator.py @@ -196,8 +196,7 @@ class OntologyGenerator: # Gửi request đến LLM result = self.llm_client.chat_json( messages=messages, - temperature=0.3, - max_tokens=4096 + temperature=0.3, ) # Kiểm tra tính hợp lệ và xử lý tinh chỉnh kết quả đầu ra @@ -206,7 +205,7 @@ class OntologyGenerator: return result # Định mức giới hạn độ dài ký tự tối đa của đoạn văn bản có thể gửi cho LLM (5 vạn chữ) - MAX_TEXT_LENGTH_FOR_LLM = 50000 + MAX_TEXT_LENGTH_FOR_LLM = 80000 def _build_user_message( self, diff --git a/backend/app/services/report_agent.py b/backend/app/services/report_agent.py index 687c0924..a390e875 100644 --- a/backend/app/services/report_agent.py +++ b/backend/app/services/report_agent.py @@ -884,6 +884,7 @@ class ReportAgent: self, graph_id: str, simulation_id: str, + project_id: str, simulation_requirement: str, llm_client: Optional[LLMClient] = None, zep_tools: Optional[ZepToolsService] = None @@ -900,10 +901,27 @@ class ReportAgent: """ self.graph_id = graph_id self.simulation_id = simulation_id + self.project_id = project_id self.simulation_requirement = simulation_requirement - self.llm = llm_client or LLMClient() - self.zep_tools = zep_tools or ZepToolsService() + self.llm = llm_client or LLMClient( + component="report_agent", + metadata={ + "simulation_id": simulation_id, + "project_id": project_id, + "phase": "report_generation", + }, + ) + self.zep_tools = zep_tools or ZepToolsService( + llm_client=LLMClient( + component="zep_tools", + metadata={ + "simulation_id": simulation_id, + "project_id": project_id, + "phase": "zep_tools", + }, + ) + ) # Định nghĩa các công cụ self.tools = self._define_tools() @@ -913,7 +931,9 @@ class ReportAgent: # Trình ghi log console (được khởi tạo trong generate_report) self.console_logger: Optional[ReportConsoleLogger] = None - logger.info(f"ReportAgent initialized: graph_id={graph_id}, simulation_id={simulation_id}") + logger.info( + f"ReportAgent initialized: graph_id={graph_id}, simulation_id={simulation_id}, project_id={project_id}" + ) def _define_tools(self) -> Dict[str, Dict[str, Any]]: """Định nghĩa các công cụ khả dụng""" @@ -1303,7 +1323,7 @@ class ReportAgent: response = self.llm.chat( messages=messages, temperature=0.5, - max_tokens=4096 + # max_tokens=4096 ) # Kiểm tra xem phản hồi có rỗng/chưa có (None) không (do API lỗi hoặc content null) @@ -1506,7 +1526,7 @@ class ReportAgent: response = self.llm.chat( messages=messages, temperature=0.5, - max_tokens=4096 + # max_tokens=4096 ) # Kiểm tra nếu ép buộc kết thúc mà LLM vẫn nhả None @@ -1561,6 +1581,19 @@ class ReportAgent: if not report_id: report_id = f"report_{uuid.uuid4().hex[:12]}" start_time = datetime.now() + + # Đồng bộ metadata chi phí LLM theo report hiện tại + # Mục tiêu: mọi call trong giai đoạn generate report đều mang report_id. + self.llm.default_metadata["report_id"] = report_id + self.llm.default_metadata["project_id"] = self.project_id + self.llm.default_metadata["simulation_id"] = self.simulation_id + self.llm.default_metadata["phase"] = "report_generation" + + zep_llm = self.zep_tools.llm + zep_llm.default_metadata["report_id"] = report_id + zep_llm.default_metadata["project_id"] = self.project_id + zep_llm.default_metadata["simulation_id"] = self.simulation_id + zep_llm.default_metadata["phase"] = "zep_tools" report = Report( report_id=report_id, diff --git a/backend/app/services/simulation_config_generator.py b/backend/app/services/simulation_config_generator.py index 926ff1b4..3b794969 100644 --- a/backend/app/services/simulation_config_generator.py +++ b/backend/app/services/simulation_config_generator.py @@ -20,6 +20,7 @@ from openai import OpenAI from ..config import Config from ..utils.logger import get_logger +from ..utils.llm_cost import create_tracked_chat_completion from .zep_entity_reader import EntityNode, ZepEntityReader logger = get_logger('mirofish.simulation_config') @@ -238,6 +239,7 @@ class SimulationConfigGenerator: api_key=self.api_key, base_url=self.base_url ) + self._runtime_metadata: Dict[str, Any] = {} def generate_config( self, @@ -269,6 +271,12 @@ class SimulationConfigGenerator: SimulationParameters: Bộ tổng cấu hình thông số đầy đủ """ logger.info(f"Start generating simulation configuration: simulation_id={simulation_id}, entity_count={len(entities)}") + self._runtime_metadata = { + "simulation_id": simulation_id, + "project_id": project_id, + "component": "simulation_config_generator", + "phase": "prepare_simulation_config", + } # Tính toán tổng số bước num_batches = math.ceil(len(entities) / self.AGENTS_PER_BATCH) @@ -438,15 +446,16 @@ class SimulationConfigGenerator: for attempt in range(max_attempts): try: - response = self.client.chat.completions.create( + response = create_tracked_chat_completion( + client=self.client, model=self.model_name, messages=[ {"role": "system", "content": system_prompt}, {"role": "user", "content": prompt} ], response_format={"type": "json_object"}, - temperature=0.7 - (attempt * 0.1) # Giảm temperature cho mỗi lần retry - # Không đặt max_tokens, cho AI sáng tạo tự do tối đa + temperature=0.7 - (attempt * 0.1), # Giảm temperature cho mỗi lần retry + metadata=self._runtime_metadata, ) content = response.choices[0].message.content diff --git a/backend/app/services/simulation_manager.py b/backend/app/services/simulation_manager.py index d9d7b022..8d6b78c3 100644 --- a/backend/app/services/simulation_manager.py +++ b/backend/app/services/simulation_manager.py @@ -342,7 +342,9 @@ class SimulationManager: graph_id=state.graph_id, # Để tìm kiếm Zep Search Index parallel_count=parallel_profile_count, # Số dòng luồng Async realtime_output_path=realtime_output_path, # Lưu log thời gian thực - output_platform=realtime_platform # Đuôi file xuất + output_platform=realtime_platform, # Đuôi file xuất + simulation_id=simulation_id, + project_id=state.project_id, ) state.profiles_count = len(profiles) diff --git a/backend/app/services/zep_tools.py b/backend/app/services/zep_tools.py index 95c1ec47..d5230d2f 100644 --- a/backend/app/services/zep_tools.py +++ b/backend/app/services/zep_tools.py @@ -436,7 +436,12 @@ class ZepToolsService: def llm(self) -> LLMClient: """Khởi tạo muộn (lazy init) cho LLM client""" if self._llm_client is None: - self._llm_client = LLMClient() + self._llm_client = LLMClient( + component="zep_tools", + metadata={ + "phase": "zep_tools", + }, + ) return self._llm_client def _call_with_retry(self, func, operation_name: str, max_retries: int = None): diff --git a/backend/app/utils/llm_client.py b/backend/app/utils/llm_client.py index faa60294..987959c5 100644 --- a/backend/app/utils/llm_client.py +++ b/backend/app/utils/llm_client.py @@ -9,6 +9,7 @@ from typing import Optional, Dict, Any, List from openai import OpenAI from ..config import Config +from .llm_cost import create_tracked_chat_completion class LLMClient: @@ -18,11 +19,15 @@ class LLMClient: self, api_key: Optional[str] = None, base_url: Optional[str] = None, - model: Optional[str] = None + model: Optional[str] = None, + component: str = "llm_client", + metadata: Optional[Dict[str, Any]] = None, ): self.api_key = api_key or Config.LLM_API_KEY self.base_url = base_url or Config.LLM_BASE_URL self.model = model or Config.LLM_MODEL_NAME + self.component = component + self.default_metadata = metadata or {} if not self.api_key: raise ValueError("LLM_API_KEY is not configured") @@ -37,7 +42,8 @@ class LLMClient: messages: List[Dict[str, str]], temperature: float = 0.7, max_tokens: int = 4096, - response_format: Optional[Dict] = None + response_format: Optional[Dict] = None, + metadata: Optional[Dict[str, Any]] = None, ) -> str: """ Gửi yêu cầu chat @@ -60,8 +66,19 @@ class LLMClient: if response_format: kwargs["response_format"] = response_format + + call_metadata = dict(self.default_metadata) + if metadata: + call_metadata.update(metadata) + call_metadata.setdefault("component", self.component) - response = self.client.chat.completions.create(**kwargs) + response = create_tracked_chat_completion( + client=self.client, + model=self.model, + messages=messages, + metadata=call_metadata, + **{k: v for k, v in kwargs.items() if k not in {"model", "messages"}}, + ) content = response.choices[0].message.content # Một số model (vd MiniMax M2.5) chèn nội dung vào content, cần loại bỏ content = re.sub(r'[\s\S]*?', '', content).strip() @@ -71,7 +88,8 @@ class LLMClient: self, messages: List[Dict[str, str]], temperature: float = 0.3, - max_tokens: int = 4096 + max_tokens: int = 4096, + metadata: Optional[Dict[str, Any]] = None, ) -> Dict[str, Any]: """ Gửi yêu cầu chat và trả về JSON @@ -88,7 +106,8 @@ class LLMClient: messages=messages, temperature=temperature, max_tokens=max_tokens, - response_format={"type": "json_object"} + response_format={"type": "json_object"}, + metadata=metadata, ) # Làm sạch markdown code fence cleaned_response = response.strip() diff --git a/backend/app/utils/llm_cost.py b/backend/app/utils/llm_cost.py new file mode 100644 index 00000000..6b2e14dc --- /dev/null +++ b/backend/app/utils/llm_cost.py @@ -0,0 +1,218 @@ +""" +Centralized LLM cost tracking utility. + +Responsibilities: +1. Calculate token cost per model. +2. Persist cost logs to logs/{project_id}/cost_{model_name}.log. +3. Persist structured JSONL records to logs/{project_id}/cost_{model_name}.jsonl. +4. Provide a single wrapped OpenAI chat-completion call so all components can reuse + one cost-accounting flow. +""" + +from __future__ import annotations + +import json +import os +import re +import threading +from datetime import datetime +from decimal import Decimal, ROUND_HALF_UP +from typing import Any, Dict, Optional + + +MODEL_COSTS_PER_1M_TOKENS: Dict[str, Dict[str, float]] = { + "Qwen/Qwen3.5-27B": {"input": 0.5, "output": 3.0}, + "gemini-3-flash": {"input": 0.5, "output": 3.0}, + "gemini-3.1-flash-lite": {"input": 0.25, "output": 1.5}, +} + + +_COUNTER_LOCK = threading.Lock() +_COUNTERS: Dict[str, Dict[str, Decimal]] = {} + + +def _decimal(value: Any) -> Decimal: + return Decimal(str(value)) + + +def _quantize_8(value: Decimal) -> Decimal: + return value.quantize(Decimal("0.00000001"), rounding=ROUND_HALF_UP) + + +def _safe_model_name(model_name: str) -> str: + if not model_name: + return "unknown_model" + safe = re.sub(r"[^A-Za-z0-9_.-]+", "_", model_name) + return safe or "unknown_model" + + +def _resolve_project_id(metadata: Optional[Dict[str, Any]]) -> str: + metadata = metadata or {} + value = metadata.get("project_id") + if value is not None and str(value).strip(): + return str(value).strip() + return "global" + + +def _resolve_component(metadata: Optional[Dict[str, Any]]) -> str: + metadata = metadata or {} + component = metadata.get("component") + if component is None: + return "unknown_component" + component = str(component).strip() + return component if component else "unknown_component" + + +def _normalize_usage(usage: Any) -> Dict[str, int]: + if usage is None: + return {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0} + + prompt_tokens = int(getattr(usage, "prompt_tokens", 0) or 0) + completion_tokens = int(getattr(usage, "completion_tokens", 0) or 0) + total_tokens = int(getattr(usage, "total_tokens", prompt_tokens + completion_tokens) or 0) + + return { + "input_tokens": prompt_tokens, + "output_tokens": completion_tokens, + "total_tokens": total_tokens, + } + + +def _resolve_model_rates(model_name: str) -> Dict[str, Decimal]: + env_in = os.environ.get("LLM_COST_INPUT_PER_1M") + env_out = os.environ.get("LLM_COST_OUTPUT_PER_1M") + if env_in is not None and env_out is not None: + return {"input": _decimal(env_in), "output": _decimal(env_out)} + + rates = MODEL_COSTS_PER_1M_TOKENS.get(model_name, {"input": 0.0, "output": 0.0}) + return {"input": _decimal(rates.get("input", 0.0)), "output": _decimal(rates.get("output", 0.0))} + + +def _get_logs_root() -> str: + # backend/app/utils -> backend/logs + return os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(__file__))), "logs") + + +def _ensure_project_log_dir(project_id: str) -> str: + root = _get_logs_root() + project_dir = os.path.join(root, project_id) + os.makedirs(project_dir, exist_ok=True) + return project_dir + + +def _build_log_paths(project_id: str, model_name: str) -> Dict[str, str]: + project_dir = _ensure_project_log_dir(project_id) + model_safe = _safe_model_name(model_name) + return { + "log": os.path.join(project_dir, f"cost_{model_safe}.log"), + "jsonl": os.path.join(project_dir, f"cost_{model_safe}.jsonl"), + } + + +def _update_counters(counter_key: str, request_cost: Decimal) -> Dict[str, Decimal]: + with _COUNTER_LOCK: + current = _COUNTERS.get(counter_key) + if current is None: + current = { + "requests": Decimal("0"), + "total_cost": Decimal("0"), + } + current["requests"] += Decimal("1") + current["total_cost"] += request_cost + _COUNTERS[counter_key] = current + return { + "requests": current["requests"], + "total_cost": current["total_cost"], + } + + +def record_llm_cost( + *, + model_name: str, + usage: Any, + metadata: Optional[Dict[str, Any]] = None, +) -> Dict[str, Any]: + """ + Calculate and persist LLM usage cost for one request. + + Returns a normalized cost payload that is also written into JSONL. + """ + metadata = metadata or {} + timestamp = datetime.now().isoformat() + project_id = _resolve_project_id(metadata) + component = _resolve_component(metadata) + usage_dict = _normalize_usage(usage) + + rates = _resolve_model_rates(model_name) + input_cost = _quantize_8(_decimal(usage_dict["input_tokens"]) * rates["input"] / _decimal(1_000_000)) + output_cost = _quantize_8(_decimal(usage_dict["output_tokens"]) * rates["output"] / _decimal(1_000_000)) + total_cost = _quantize_8(input_cost + output_cost) + + record = { + "timestamp": timestamp, + "model": model_name, + "input_tokens": usage_dict["input_tokens"], + "output_tokens": usage_dict["output_tokens"], + "total_tokens": usage_dict["total_tokens"], + "input_cost_usd": float(input_cost), + "output_cost_usd": float(output_cost), + "total_cost_usd": float(total_cost), + "metadata": { + "component": component, + "simulation_id": metadata.get("simulation_id"), + "platform": metadata.get("platform"), + "phase": metadata.get("phase"), + "project_id": project_id, + "report_id": metadata.get("report_id"), + }, + } + + paths = _build_log_paths(project_id=project_id, model_name=model_name) + counter_key = f"{project_id}::{model_name}" + counter = _update_counters(counter_key, total_cost) + request_no = int(counter["requests"]) + cumulative_cost = _quantize_8(counter["total_cost"]) + + line = ( + f"[{timestamp}] [Request {request_no}] [{component}] " + f"Called model: {model_name}, " + f"input_tokens: {usage_dict['input_tokens']} | " + f"output_tokens: {usage_dict['output_tokens']} | " + f"total_tokens: {usage_dict['total_tokens']} | " + f"input_cost_usd: {float(input_cost):.8f} | " + f"output_cost_usd: {float(output_cost):.8f} | " + f"total_cost_usd: {float(total_cost):.8f} | " + f"cumulative_total_cost_usd: {float(cumulative_cost):.8f}" + ) + + with open(paths["log"], "a", encoding="utf-8") as log_f: + log_f.write(line + "\n") + + with open(paths["jsonl"], "a", encoding="utf-8") as jsonl_f: + jsonl_f.write(json.dumps(record, ensure_ascii=False) + "\n") + + return record + + +def create_tracked_chat_completion( + *, + client: Any, + model: str, + messages: Any, + metadata: Optional[Dict[str, Any]] = None, + **kwargs: Any, +) -> Any: + """ + Single entry-point for OpenAI-compatible chat completion + cost logging. + + All modules should call this wrapper instead of calling + client.chat.completions.create directly. + """ + response = client.chat.completions.create( + model=model, + messages=messages, + **kwargs, + ) + record_llm_cost(model_name=model, usage=getattr(response, "usage", None), metadata=metadata) + return response + diff --git a/backend/scripts/llm_cost_patch.py b/backend/scripts/llm_cost_patch.py new file mode 100644 index 00000000..b02cde31 --- /dev/null +++ b/backend/scripts/llm_cost_patch.py @@ -0,0 +1,87 @@ +"""Patch OpenAI chat completion calls in simulation scripts for centralized cost logging.""" + +from __future__ import annotations + +import copy +from typing import Any, Dict, Optional + +from app.utils.llm_cost import record_llm_cost + + +_PATCHED = False +_PATCH_CONTEXT: Dict[str, Any] = {} + + +def _build_metadata(model: str) -> Dict[str, Any]: + metadata = copy.deepcopy(_PATCH_CONTEXT) + metadata.setdefault("component", "scripts.simulation") + metadata.setdefault("phase", "simulation_run") + metadata.setdefault("model", model) + return metadata + + +def install_openai_cost_patch( + *, + simulation_id: Optional[str], + project_id: Optional[str], + platform: str, + component: str, + phase: str = "simulation_run", +) -> None: + """ + Install monkey-patch for OpenAI SDK calls used indirectly by CAMEL/OASIS. + + The patch is process-scoped and idempotent. + """ + global _PATCHED + global _PATCH_CONTEXT + + _PATCH_CONTEXT = { + "simulation_id": simulation_id, + "project_id": project_id, + "platform": platform, + "component": component, + "phase": phase, + } + + if _PATCHED: + return + + try: + from openai.resources.chat.completions.completions import Completions, AsyncCompletions + except Exception as exc: # pragma: no cover + print(f"[llm_cost_patch] Failed to import OpenAI completion classes: {exc}") + return + + original_sync_create = Completions.create + original_async_create = AsyncCompletions.create + + def sync_create_wrapper(self, *args, **kwargs): + model = kwargs.get("model", "unknown_model") + response = original_sync_create(self, *args, **kwargs) + try: + record_llm_cost( + model_name=model, + usage=getattr(response, "usage", None), + metadata=_build_metadata(model), + ) + except Exception as exc: + print(f"[llm_cost_patch] Failed to record sync cost: {exc}") + return response + + async def async_create_wrapper(self, *args, **kwargs): + model = kwargs.get("model", "unknown_model") + response = await original_async_create(self, *args, **kwargs) + try: + record_llm_cost( + model_name=model, + usage=getattr(response, "usage", None), + metadata=_build_metadata(model), + ) + except Exception as exc: + print(f"[llm_cost_patch] Failed to record async cost: {exc}") + return response + + Completions.create = sync_create_wrapper + AsyncCompletions.create = async_create_wrapper + _PATCHED = True diff --git a/backend/scripts/run_parallel_simulation.py b/backend/scripts/run_parallel_simulation.py index 4ccde2d5..689fc103 100644 --- a/backend/scripts/run_parallel_simulation.py +++ b/backend/scripts/run_parallel_simulation.py @@ -156,6 +156,7 @@ def init_logging_for_simulation(simulation_dir: str): from action_logger import SimulationLogManager, PlatformActionLogger +from llm_cost_patch import install_openai_cost_patch try: from camel.models import ModelFactory @@ -1533,6 +1534,14 @@ async def main(): config = load_config(args.config) simulation_dir = os.path.dirname(args.config) or "." wait_for_commands = not args.no_wait + + install_openai_cost_patch( + simulation_id=config.get("simulation_id"), + project_id=config.get("project_id"), + platform="parallel", + component="scripts.run_parallel_simulation", + phase="simulation_run", + ) # Khởi tạo cấu hình log (tắt log OASIS, dọn file cũ) init_logging_for_simulation(simulation_dir) diff --git a/backend/scripts/run_reddit_simulation.py b/backend/scripts/run_reddit_simulation.py index cc063616..897a75fb 100644 --- a/backend/scripts/run_reddit_simulation.py +++ b/backend/scripts/run_reddit_simulation.py @@ -48,6 +48,7 @@ else: import re +from llm_cost_patch import install_openai_cost_patch class UnicodeFormatter(logging.Formatter): @@ -731,6 +732,15 @@ async def main(): config_path=args.config, wait_for_commands=not args.no_wait ) + + install_openai_cost_patch( + simulation_id=runner.config.get("simulation_id"), + project_id=runner.config.get("project_id"), + platform="reddit", + component="scripts.run_reddit_simulation", + phase="simulation_run", + ) + await runner.run(max_rounds=args.max_rounds) diff --git a/backend/scripts/run_twitter_simulation.py b/backend/scripts/run_twitter_simulation.py index aa3cec53..244dcac5 100644 --- a/backend/scripts/run_twitter_simulation.py +++ b/backend/scripts/run_twitter_simulation.py @@ -48,6 +48,7 @@ else: import re +from llm_cost_patch import install_openai_cost_patch class UnicodeFormatter(logging.Formatter): @@ -743,6 +744,15 @@ async def main(): config_path=args.config, wait_for_commands=not args.no_wait ) + + install_openai_cost_patch( + simulation_id=runner.config.get("simulation_id"), + project_id=runner.config.get("project_id"), + platform="twitter", + component="scripts.run_twitter_simulation", + phase="simulation_run", + ) + await runner.run(max_rounds=args.max_rounds)