CoolFace
Apppublic

Rhinox13/chatapi

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
pending.py366 linesDownload Raw Back to services
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