Skip to content

Commit 10b4918

Browse files
committed
Update antigravity.py
1 parent 6a39565 commit 10b4918

1 file changed

Lines changed: 14 additions & 133 deletions

File tree

src/api/antigravity.py

Lines changed: 14 additions & 133 deletions
Original file line numberDiff line numberDiff line change
@@ -7,10 +7,7 @@
77
import copy
88
import hashlib
99
import json
10-
import os
11-
import time
1210
import uuid
13-
from dataclasses import dataclass
1411
from datetime import datetime, timezone
1512
from typing import Any, Dict, List, Optional, Callable, Tuple
1613

@@ -42,51 +39,6 @@
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-
9042
def _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

18662
def _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

Comments
 (0)