77import copy
88import hashlib
99import json
10- import os
11- import time
1210import uuid
13- from dataclasses import dataclass
1411from datetime import datetime , timezone
1512from typing import Any , Dict , List , Optional , Callable , Tuple
1613
4239# 使用全局单例 credential_manager,自动初始化
4340
4441
45- # ==================== 会话状态管理 ====================
46-
47- SESSION_TTL_SECONDS = 6 * 60 * 60
48- MAX_SESSION_STATES = 1024
49- _REDIS_KEY_PREFIX = "antigravity:session:"
50-
51-
52- @dataclass
53- class AntigravitySessionState :
54- conversation_id : str
55- trajectory_id : str
56- session_id : str
57- step_index : int
58- created_at : float
59- last_used_at : float
60-
61-
62- # 内存回退存储
63- _session_states : Dict [str , AntigravitySessionState ] = {}
64-
65- # Redis 客户端(懒初始化,REDIS_URL 存在时使用)
66- _redis_client = None
67- _redis_checked = False
68-
69-
70- async def _get_redis ():
71- """懒初始化 Redis 客户端,REDIS_URL 未设置时返回 None。"""
72- global _redis_client , _redis_checked
73- if _redis_checked :
74- return _redis_client
75- _redis_checked = True
76- redis_url = os .getenv ("REDIS_URL" )
77- if not redis_url :
78- return None
79- try :
80- import redis .asyncio as aioredis # type: ignore
81- client = aioredis .from_url (redis_url , decode_responses = True )
82- await client .ping ()
83- _redis_client = client
84- log .info ("[SESSION] Redis session store enabled" )
85- except Exception as e :
86- log .warning (f"[SESSION] Redis unavailable, falling back to in-memory: { e } " )
87- return _redis_client
88-
89-
9042def _extract_first_user_text (request_payload : Dict [str , Any ]) -> str :
9143 contents = request_payload .get ("contents" , [])
9244 if not isinstance (contents , list ):
@@ -103,84 +55,8 @@ def _extract_first_user_text(request_payload: Dict[str, Any]) -> str:
10355 return ""
10456
10557
106- def _session_key (request_payload : Dict [str , Any ], model : str = "" ) -> str :
107- session_id = request_payload .get ("sessionId" )
108- if session_id :
109- return f"session:{ session_id } "
110- model_prefix = f"model:{ model } :" if model else ""
111- first_user_text = _extract_first_user_text (request_payload )
112- if first_user_text :
113- digest = hashlib .sha256 (first_user_text .encode ("utf-8" )).hexdigest ()[:32 ]
114- return f"{ model_prefix } text:{ digest } "
115- return f"{ model_prefix } default"
116-
117-
118- def _prune_session_states (now : float ) -> None :
119- expired = [k for k , s in _session_states .items () if now - s .last_used_at > SESSION_TTL_SECONDS ]
120- for k in expired :
121- _session_states .pop (k , None )
122- if len (_session_states ) <= MAX_SESSION_STATES :
123- return
124- overflow = len (_session_states ) - MAX_SESSION_STATES
125- oldest = sorted (_session_states .items (), key = lambda item : item [1 ].last_used_at )
126- for k , _ in oldest [:overflow ]:
127- _session_states .pop (k , None )
128-
129-
130- def _make_new_state (first_user_text : str , now : float ) -> AntigravitySessionState :
131- if first_user_text :
132- digest = hashlib .sha256 (first_user_text .encode ("utf-8" )).digest ()
133- session_id_val = int .from_bytes (digest [:8 ], "big" ) & 0x7FFFFFFFFFFFFFFF
134- session_id = f"-{ session_id_val } "
135- else :
136- session_id = f"-{ uuid .uuid4 ().int % 9_000_000_000_000_000_000 } "
137- return AntigravitySessionState (
138- conversation_id = str (uuid .uuid4 ()),
139- trajectory_id = str (uuid .uuid4 ()),
140- session_id = session_id ,
141- step_index = 1 ,
142- created_at = now ,
143- last_used_at = now ,
144- )
145-
146-
147- async def _get_session_state (request_payload : Dict [str , Any ], model : str = "" ) -> AntigravitySessionState :
148- now = time .time ()
149- key = _session_key (request_payload , model )
150- first_user_text = _extract_first_user_text (request_payload )
151-
152- redis = await _get_redis ()
153- if redis is not None :
154- redis_key = f"{ _REDIS_KEY_PREFIX } { key } "
155- try :
156- raw = await redis .get (redis_key )
157- if raw :
158- data = json .loads (raw )
159- state = AntigravitySessionState (** data )
160- state .step_index += 1
161- state .last_used_at = now
162- else :
163- state = _make_new_state (first_user_text , now )
164- await redis .set (redis_key , json .dumps (state .__dict__ ), ex = SESSION_TTL_SECONDS )
165- return state
166- except Exception as e :
167- log .warning (f"[SESSION] Redis error, falling back to memory: { e } " )
168-
169- # 内存回退
170- _prune_session_states (now )
171- state = _session_states .get (key )
172- if state :
173- state .step_index += 1
174- state .last_used_at = now
175- return state
176- state = _make_new_state (first_user_text , now )
177- _session_states [key ] = state
178- return state
179-
180-
181- def _generate_request_id (conversation_id : str , trajectory_id : str , step : int ) -> str :
182- unix_ms = int (datetime .now (timezone .utc ).timestamp () * 1000 )
183- return f"agent/{ conversation_id } /{ unix_ms } /{ trajectory_id } /{ step } "
58+ def _generate_request_id () -> str :
59+ return f"agent/{ uuid .uuid4 ()} "
18460
18561
18662def _build_labels (model : str , trajectory_id : str , step : int ) -> Dict [str , str ]:
@@ -232,19 +108,24 @@ async def wrap_cli_request(
232108 返回 (payload, request_id)。
233109 """
234110 inner = copy .deepcopy (gemini_request )
111+ first_user_text = _extract_first_user_text (inner )
235112
236113 # 移除 safetySettings(CLI 不发送)
237114 inner .pop ("safetySettings" , None )
238115
239- # 获取/更新会话状态
240- state = await _get_session_state (inner , model )
241-
242116 # 注入 sessionId
243- if not inner .get ("sessionId" ):
244- inner ["sessionId" ] = state .session_id
117+ session_id = str (inner .get ("sessionId" ) or "" ).strip ()
118+ if not session_id :
119+ if first_user_text :
120+ digest = hashlib .sha256 (first_user_text .encode ("utf-8" )).digest ()
121+ session_id_val = int .from_bytes (digest [:8 ], "big" ) & 0x7FFFFFFFFFFFFFFF
122+ session_id = f"-{ session_id_val } "
123+ else :
124+ session_id = f"-{ uuid .uuid4 ().int % 9_000_000_000_000_000_000 } "
125+ inner ["sessionId" ] = session_id
245126
246127 # 注入 labels
247- inner ["labels" ] = _build_labels (model , state . trajectory_id , state . step_index )
128+ inner ["labels" ] = _build_labels (model , session_id , 1 )
248129
249130 # toolConfig 默认 VALIDATED
250131 tool_config = inner .get ("toolConfig" ) or {}
@@ -253,7 +134,7 @@ async def wrap_cli_request(
253134 tool_config ["functionCallingConfig" ] = func_config
254135 inner ["toolConfig" ] = tool_config
255136
256- request_id = _generate_request_id (state . conversation_id , state . trajectory_id , state . step_index )
137+ request_id = _generate_request_id ()
257138
258139 payload = {
259140 "project" : project_id ,
0 commit comments