Rhinox13/chatapi
0
1from __future__ import annotations
2
3import json
4import threading
5import time
6import uuid
7from dataclasses import dataclass, field
8from typing import Any
9
10
11@dataclass
12class PendingTurn:
13 request_id: str
14 conversation_id: str
15 owner_id: str
16 model: str
17 input_text: str
18 request_format: str = "responses"
19 reasoning_stream_mode: str = ""
20 created_at: float = field(default_factory=time.time)
21 max_age_seconds: float = 0.0
22 auto_abort_message: str = ""
23 max_output_chars: int = 0
24 output_limit_abort_message: str = ""
25 event: threading.Event = field(default_factory=threading.Event)
26 stream_event: threading.Event = field(default_factory=threading.Event)
27 assistant_text: str = ""
28 response_id: str = ""
29 draft_chunks: list[Any] = field(default_factory=list)
30 draft_text: str = ""
31 draft_segments: list[dict[str, str]] = field(default_factory=list)
32 draft_answer_text: str = ""
33 output_chars: int = 0
34 aborted: bool = False
35 abort_message: str = ""
36 resolved: bool = False
37 response_mode: str = "assistant_message"
38 response_output_items: list[dict[str, Any]] = field(default_factory=list)
39 response_output_text: str = ""
40 available_tool_names: set[str] = field(default_factory=set)
41 available_tool_schemas: dict[str, dict[str, Any]] = field(default_factory=dict)
42
43
44class PendingTurnRegistry:
45 def __init__(self) -> None:
46 self._lock = threading.Lock()
47 self._by_request_id: dict[str, PendingTurn] = {}
48 self._by_conversation_id: dict[str, str] = {}
49
50 def register(
51 self,
52 *,
53 conversation_id: str,
54 owner_id: str,
55 model: str,
56 input_text: str,
57 request_format: str = "responses",
58 reasoning_stream_mode: str = "",
59 max_age_seconds: float = 0.0,
60 auto_abort_message: str = "",
61 max_output_chars: int = 0,
62 output_limit_abort_message: str = "",
63 available_tool_names: set[str] | None = None,
64 available_tool_schemas: dict[str, dict[str, Any]] | None = None,
65 ) -> PendingTurn:
66 with self._lock:
67 if conversation_id in self._by_conversation_id:
68 existing_request_id = self._by_conversation_id.get(conversation_id)
69 existing = self._by_request_id.get(existing_request_id or "")
70 if existing is not None and existing.event.is_set():
71 self._by_conversation_id.pop(conversation_id, None)
72 else:
73 raise ValueError("conversation is waiting for a reply")
74 pending = PendingTurn(
75 request_id=f"resp_{uuid.uuid4().hex}",
76 conversation_id=conversation_id,
77 owner_id=owner_id,
78 model=model,
79 input_text=input_text,
80 request_format=request_format,
81 reasoning_stream_mode=reasoning_stream_mode,
82 max_age_seconds=max(0.0, float(max_age_seconds or 0.0)),
83 auto_abort_message=str(auto_abort_message or ""),
84 max_output_chars=max(0, int(max_output_chars or 0)),
85 output_limit_abort_message=str(output_limit_abort_message or ""),
86 available_tool_names=available_tool_names or set(),
87 available_tool_schemas=available_tool_schemas or {},
88 )
89 self._by_request_id[pending.request_id] = pending
90 self._by_conversation_id[conversation_id] = pending.request_id
91 return pending
92
93 def get_by_conversation(self, conversation_id: str) -> PendingTurn | None:
94 with self._lock:
95 request_id = self._by_conversation_id.get(conversation_id)
96 if not request_id:
97 return None
98 pending = self._by_request_id.get(request_id)
99 if pending is not None and pending.event.is_set():
100 self._by_conversation_id.pop(conversation_id, None)
101 return None
102 return pending
103
104 def active_count_by_owner(self, owner_id: str) -> int:
105 with self._lock:
106 return sum(
107 1
108 for pending in self._by_request_id.values()
109 if pending.owner_id == owner_id and not pending.event.is_set()
110 )
111
112 def abort_owner_over_limit(
113 self,
114 *,
115 owner_id: str,
116 max_active: int,
117 error_message: str,
118 ) -> list[PendingTurn]:
119 if max_active <= 0:
120 return []
121 aborted: list[PendingTurn] = []
122 with self._lock:
123 active = sorted(
124 (
125 pending
126 for pending in self._by_request_id.values()
127 if pending.owner_id == owner_id and not pending.event.is_set()
128 ),
129 key=lambda pending: pending.created_at,
130 )
131 while len(active) >= max_active:
132 pending = active.pop(0)
133 aborted.append(self._mark_aborted_locked(pending, error_message))
134 return aborted
135
136 def abort_expired(self, *, max_age_seconds: float, error_message: str) -> list[PendingTurn]:
137 if max_age_seconds <= 0:
138 return []
139 deadline = time.time() - max_age_seconds
140 aborted: list[PendingTurn] = []
141 with self._lock:
142 for pending in list(self._by_request_id.values()):
143 if pending.event.is_set() or pending.created_at > deadline:
144 continue
145 aborted.append(self._mark_aborted_locked(pending, error_message))
146 return aborted
147
148 def resolve(
149 self,
150 *,
151 conversation_id: str,
152 owner_id: str,
153 assistant_text: str,
154 response_id: str,
155 response_mode: str = "assistant_message",
156 response_output_items: list[dict[str, Any]] | None = None,
157 response_output_text: str | None = None,
158 ) -> PendingTurn:
159 with self._lock:
160 request_id = self._by_conversation_id.get(conversation_id)
161 if not request_id:
162 raise ValueError("conversation is not waiting for a reply")
163 pending = self._by_request_id.get(request_id)
164 if pending is None or pending.owner_id != owner_id:
165 raise ValueError("conversation is not waiting for a reply")
166 if pending.resolved or pending.event.is_set():
167 raise ValueError("conversation reply is already completed")
168 pending.assistant_text = assistant_text
169 pending.response_id = response_id
170 pending.response_mode = response_mode
171 pending.response_output_items = list(response_output_items or [])
172 pending.response_output_text = (
173 assistant_text if response_output_text is None else response_output_text
174 )
175 pending.resolved = True
176 pending.stream_event.set()
177 pending.event.set()
178 return pending
179
180 def consume_draft_chunks(self, request_id: str) -> list[Any]:
181 with self._lock:
182 pending = self._by_request_id.get(request_id)
183 if pending is None:
184 return []
185 chunks = list(pending.draft_chunks)
186 pending.draft_chunks.clear()
187 return chunks
188
189 def add_draft(
190 self,
191 *,
192 conversation_id: str,
193 owner_id: str,
194 chunk: str,
195 kind: str = "answer",
196 ) -> PendingTurn:
197 with self._lock:
198 request_id = self._by_conversation_id.get(conversation_id)
199 if not request_id:
200 raise ValueError("conversation is not waiting for a reply")
201 pending = self._by_request_id.get(request_id)
202 if pending is None or pending.owner_id != owner_id:
203 raise ValueError("conversation is not waiting for a reply")
204 if pending.resolved or pending.event.is_set():
205 raise ValueError("conversation reply is already completed")
206 chunk = str(chunk or "")
207 output_limit_abort = self._mark_aborted_if_output_limit_exceeded_locked(
208 pending,
209 chunk,
210 )
211 if output_limit_abort is not None:
212 return output_limit_abort
213 normalized_kind = "thinking" if str(kind).strip() == "thinking" else "answer"
214 if normalized_kind == "thinking":
215 if not pending.draft_segments and pending.draft_text:
216 pending.draft_segments.append({"type": "answer", "text": pending.draft_text})
217 if pending.draft_segments and pending.draft_segments[-1]["type"] == "thinking":
218 pending.draft_segments[-1]["text"] += chunk
219 else:
220 pending.draft_segments.append({"type": "thinking", "text": chunk})
221 pending.draft_chunks.append({"kind": "thinking", "text": chunk})
222 pending.draft_text = _serialize_draft_segments(pending.draft_segments)
223 elif pending.draft_segments:
224 if pending.draft_segments[-1]["type"] == "answer":
225 pending.draft_segments[-1]["text"] += chunk
226 else:
227 pending.draft_segments.append({"type": "answer", "text": chunk})
228 pending.draft_chunks.append({"kind": "answer", "text": chunk})
229 pending.draft_answer_text += chunk
230 pending.draft_text = _serialize_draft_segments(pending.draft_segments)
231 else:
232 pending.draft_chunks.append(chunk)
233 pending.draft_text += chunk
234 pending.draft_answer_text += chunk
235 pending.output_chars += len(chunk)
236 pending.stream_event.set()
237 return pending
238
239 def abort(
240 self,
241 *,
242 conversation_id: str,
243 owner_id: str,
244 error_message: str,
245 ) -> PendingTurn:
246 with self._lock:
247 request_id = self._by_conversation_id.get(conversation_id)
248 if not request_id:
249 raise ValueError("conversation is not waiting for a reply")
250 pending = self._by_request_id.get(request_id)
251 if pending is None or pending.owner_id != owner_id:
252 raise ValueError("conversation is not waiting for a reply")
253 if pending.resolved or pending.event.is_set():
254 raise ValueError("conversation reply is already completed")
255 return self._mark_aborted_locked(pending, error_message)
256
257 def abort_by_request_id(self, *, request_id: str, error_message: str) -> PendingTurn | None:
258 with self._lock:
259 pending = self._by_request_id.get(request_id)
260 if pending is None or pending.resolved or pending.event.is_set():
261 return None
262 return self._mark_aborted_locked(pending, error_message)
263
264 def abort_if_expired(self, request_id: str) -> PendingTurn | None:
265 with self._lock:
266 pending = self._by_request_id.get(request_id)
267 if pending is None or pending.resolved or pending.event.is_set():
268 return None
269 if pending.max_age_seconds <= 0:
270 return None
271 if time.time() - pending.created_at < pending.max_age_seconds:
272 return None
273 return self._mark_aborted_locked(
274 pending,
275 pending.auto_abort_message or "本次回复等待超过限制,已自动结束,请重新发送。",
276 )
277
278 def abort_if_output_would_exceed(
279 self,
280 *,
281 request_id: str,
282 extra_text: str,
283 ) -> PendingTurn | None:
284 with self._lock:
285 pending = self._by_request_id.get(request_id)
286 if pending is None or pending.resolved or pending.event.is_set():
287 return None
288 return self._mark_aborted_if_output_limit_exceeded_locked(
289 pending,
290 str(extra_text or ""),
291 )
292
293 @staticmethod
294 def _clear_draft_locked(pending: PendingTurn) -> None:
295 pending.draft_chunks.clear()
296 pending.draft_text = ""
297 pending.draft_segments.clear()
298 pending.draft_answer_text = ""
299 pending.output_chars = 0
300
301 def _mark_aborted_if_output_limit_exceeded_locked(
302 self,
303 pending: PendingTurn,
304 extra_text: str,
305 ) -> PendingTurn | None:
306 if pending.max_output_chars <= 0:
307 return None
308 if not extra_text:
309 return None
310 if pending.output_chars + len(extra_text) <= pending.max_output_chars:
311 return None
312 return self._mark_aborted_locked(
313 pending,
314 pending.output_limit_abort_message or "本次回复超过长度限制,已自动结束,请重新发送。",
315 )
316
317 def _mark_aborted_locked(self, pending: PendingTurn, error_message: str) -> PendingTurn:
318 pending.aborted = True
319 pending.abort_message = error_message
320 self._clear_draft_locked(pending)
321 pending.stream_event.set()
322 pending.event.set()
323 return pending
324
325 def discard(
326 self,
327 *,
328 conversation_id: str,
329 owner_id: str,
330 ) -> PendingTurn | None:
331 with self._lock:
332 request_id = self._by_conversation_id.get(conversation_id)
333 if not request_id:
334 return None
335 pending = self._by_request_id.get(request_id)
336 if pending is None or pending.owner_id != owner_id:
337 return None
338 self._by_request_id.pop(request_id, None)
339 self._by_conversation_id.pop(conversation_id, None)
340 return pending
341
342 def wait(self, request_id: str) -> PendingTurn:
343 with self._lock:
344 pending = self._by_request_id.get(request_id)
345 if pending is None:
346 raise ValueError("response not found")
347 pending.event.wait()
348 with self._lock:
349 self._by_request_id.pop(request_id, None)
350 self._by_conversation_id.pop(pending.conversation_id, None)
351 return pending
352
353
354def _serialize_draft_segments(segments: list[dict[str, str]]) -> str:
355 payload: list[dict[str, str]] = []
356 for segment in segments:
357 text = str(segment.get("text") or "")
358 if not text:
359 continue
360 segment_type = "reasoning_text" if segment.get("type") == "thinking" else "output_text"
361 payload.append({
362 "type": segment_type,
363 "text": text,
364 })
365 return json.dumps(payload, ensure_ascii=False)
366 