/
man4j
/
agent-server
Обзор
Документация
Войти
/
man4j
/
agent-server
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
src/agent_server/message_window.py
299 строк
10 KB
Vladimir
fixes
28 апр 2026, 19:57
28 апр 2026, 19:57
f04f1da
Код
Авторство
О чём код?
import json from agent_server.chat.types import ChatMessage from agent_server.config.env import DEFAULT_CHAR_PER_TOKEN from agent_server.history_budget import clamp_recent_history_ratio CHAT_SUMMARY_MESSAGE_NAME = "chat_summary" class ContextWindowExceededError(ValueError): pass def estimate_tokens_for_text(text: str, char_per_token: float) -> int: if not text: return 0 return max(1, int(len(text) / max(1.0, char_per_token))) def estimate_tokens_for_message(message: ChatMessage, char_per_token: float) -> int: total = 6 total += estimate_tokens_for_text(message.get("content") or "", char_per_token) total += estimate_tokens_for_text(message.get("name") or "", char_per_token) total += estimate_tokens_for_text(message.get("reasoning_content") or "", char_per_token) if message.get("role") == "tool": total += estimate_tokens_for_text(message.get("tool_call_id") or "", char_per_token) if message.get("tool_calls"): try: total += estimate_tokens_for_text( json.dumps(message["tool_calls"], ensure_ascii=False), char_per_token, ) except Exception: total += 50 return total def estimated_prompt_tokens(messages: list[ChatMessage], char_per_token: float) -> int: return sum(estimate_tokens_for_message(message, char_per_token) for message in messages) def clear_reasoning_content(messages: list[ChatMessage]) -> None: for message in messages: if message.get("role") == "assistant": message.pop("reasoning_content", None) def split_leading_system_messages( messages: list[ChatMessage], ) -> tuple[list[ChatMessage], list[ChatMessage]]: prefix: list[ChatMessage] = [] idx = 0 summary_message: ChatMessage | None = None while idx < len(messages) and messages[idx].get("role") == "system": message = messages[idx] if message.get("name") == CHAT_SUMMARY_MESSAGE_NAME: summary_message = message else: prefix.append(message) idx += 1 if summary_message is not None: if prefix: prefix.insert(1, summary_message) else: prefix.append(summary_message) return prefix, messages[idx:] def _split_into_turns(messages: list[ChatMessage]) -> list[list[ChatMessage]]: turns: list[list[ChatMessage]] = [] current_turn: list[ChatMessage] = [] for message in messages: if message.get("role") == "user": if current_turn: turns.append(current_turn) current_turn = [message] else: if current_turn: current_turn.append(message) else: # хвост без user в начале сохраняем отдельным "turn", # чтобы ничего не потерять при восстановлении истории turns.append([message]) if current_turn: turns.append(current_turn) return turns def _estimate_turn_tokens(turn: list[ChatMessage], char_per_token: float) -> int: return sum(estimate_tokens_for_message(message, char_per_token) for message in turn) def _tool_call_ids(message: ChatMessage) -> set[str]: ids: set[str] = set() for tool_call in message.get("tool_calls") or []: tool_call_id = tool_call.get("id") if tool_call_id: ids.add(tool_call_id) return ids def _split_turn_into_trim_units(turn: list[ChatMessage]) -> list[list[ChatMessage]]: units: list[list[ChatMessage]] = [] idx = 0 while idx < len(turn): message = turn[idx] tool_call_ids = _tool_call_ids(message) if message.get("role") == "assistant" and tool_call_ids: unit = [message] idx += 1 while ( idx < len(turn) and turn[idx].get("role") == "tool" and turn[idx].get("tool_call_id") in tool_call_ids ): unit.append(turn[idx]) idx += 1 units.append(unit) continue units.append([message]) idx += 1 return units def _trim_single_turn_to_fit( system_prefix: list[ChatMessage], turn: list[ChatMessage], char_per_token: float, history_token_limit: int, ) -> list[ChatMessage]: system_tokens = _ensure_system_prefix_fits( system_prefix, char_per_token, history_token_limit, ) units = _split_turn_into_trim_units(turn) protected_prefix: list[list[ChatMessage]] = [] removable_units = units if units and any(message.get("role") == "user" for message in units[0]): protected_prefix = [units[0]] removable_units = units[1:] protected_tokens = sum( _estimate_turn_tokens(unit, char_per_token) for unit in protected_prefix ) removable_tokens = [ _estimate_turn_tokens(unit, char_per_token) for unit in removable_units ] total_tokens = system_tokens + protected_tokens + sum(removable_tokens) first_removable_idx = 0 while first_removable_idx < len(removable_units) and total_tokens > history_token_limit: total_tokens -= removable_tokens[first_removable_idx] first_removable_idx += 1 if total_tokens > history_token_limit: if protected_prefix: raise ContextWindowExceededError( "Последний пользовательский запрос не помещается в контекст даже после подрезания истории." ) return [] kept_units = removable_units[first_removable_idx:] return [message for unit in protected_prefix + kept_units for message in unit] def _ensure_system_prefix_fits( system_prefix: list[ChatMessage], char_per_token: float, history_token_limit: int, ) -> int: system_tokens = estimated_prompt_tokens(system_prefix, char_per_token) if system_tokens > history_token_limit: raise ContextWindowExceededError( "Системный промпт и summary не помещаются в доступный бюджет истории." ) return system_tokens def compute_trimmed_history( messages: list[ChatMessage], char_per_token: float, context_size: int, response_headroom_tokens: int = 32000, recent_history_ratio: float = 0.5, ) -> tuple[list[ChatMessage], list[ChatMessage]]: """ Returns a tuple: (trimmed_messages, dropped_messages). Keeps all leading system messages intact. Then preserves an adaptive recent conversational tail whose size is limited by recent_history_ratio * history budget. After that, if the full prompt still does not fit, trims older turns until it fits. """ if not messages: return [], [] system_prefix, rest = split_leading_system_messages(messages) rest = list(rest) original_rest = list(rest) history_token_limit = max(1, context_size - response_headroom_tokens) recent_history_ratio = clamp_recent_history_ratio(recent_history_ratio) recent_budget = max(1, int(history_token_limit * recent_history_ratio)) system_tokens = _ensure_system_prefix_fits( system_prefix, char_per_token, history_token_limit, ) kept_turns: list[list[ChatMessage]] = [] kept_tokens = 0 if rest: turns = _split_into_turns(rest) turn_entries = [ (turn, _estimate_turn_tokens(turn, char_per_token)) for turn in turns ] kept_turns_from_end: list[list[ChatMessage]] = [] used_tokens = 0 for turn, turn_tokens in reversed(turn_entries): if kept_turns_from_end and used_tokens + turn_tokens > recent_budget: break kept_turns_from_end.append(turn) used_tokens += turn_tokens kept_turns = list(reversed(kept_turns_from_end)) kept_tokens = used_tokens first_kept_turn_idx = 0 kept_turn_tokens = [_estimate_turn_tokens(turn, char_per_token) for turn in kept_turns] while ( first_kept_turn_idx < len(kept_turns) and system_tokens + kept_tokens > history_token_limit ): if len(kept_turns) - first_kept_turn_idx <= 1: kept_turns = [ _trim_single_turn_to_fit( system_prefix, kept_turns[first_kept_turn_idx], char_per_token, history_token_limit, ) ] first_kept_turn_idx = 0 break kept_tokens -= kept_turn_tokens[first_kept_turn_idx] first_kept_turn_idx += 1 rest = [message for turn in kept_turns[first_kept_turn_idx:] for message in turn] if rest and estimated_prompt_tokens(system_prefix + rest, char_per_token) > history_token_limit: turns = _split_into_turns(rest) rest = _trim_single_turn_to_fit( system_prefix, turns[-1] if turns else [], char_per_token, history_token_limit, ) kept_message_ids = {id(message) for message in rest} dropped = [message for message in original_rest if id(message) not in kept_message_ids] return system_prefix + rest, dropped def trim_history_by_budget( messages: list[ChatMessage], char_per_token: float, context_size: int, response_headroom_tokens: int = 32000, recent_history_ratio: float = 0.5, ) -> None: trimmed, _ = compute_trimmed_history( messages, char_per_token, context_size, response_headroom_tokens=response_headroom_tokens, recent_history_ratio=recent_history_ratio, ) messages[:] = trimmed