263 lines
7.9 KiB
Python
263 lines
7.9 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import logging
|
|
|
|
from yuxi.channel.extensions.feishu.card_kit import (
|
|
CardKitState,
|
|
open_card_stream,
|
|
push_card_update,
|
|
close_card_stream,
|
|
send_card_fallback,
|
|
should_send_update,
|
|
merge_content,
|
|
STREAMING_THROTTLE_MS,
|
|
SIGNIFICANT_DELTA_CHARS,
|
|
)
|
|
|
|
from yuxi.channel.extensions.feishu.client import get_client
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
class FeishuStreamingAdapter:
|
|
def __init__(
|
|
self,
|
|
throttle_ms: int = STREAMING_THROTTLE_MS,
|
|
significant_delta_chars: int = SIGNIFICANT_DELTA_CHARS,
|
|
):
|
|
self._throttle_ms = throttle_ms
|
|
self._significant_delta_chars = significant_delta_chars
|
|
|
|
async def start_stream(
|
|
self,
|
|
app_id: str,
|
|
app_secret: str,
|
|
domain: str,
|
|
target_id: str,
|
|
reply_to_id: str | None = None,
|
|
thread_id: str | None = None,
|
|
http_timeout_ms: int = 30000,
|
|
) -> CardKitState | None:
|
|
try:
|
|
client = get_client(app_id, app_secret, domain, http_timeout_ms)
|
|
except Exception:
|
|
logger.exception("Failed to create client for streaming")
|
|
return None
|
|
|
|
return await open_card_stream(
|
|
client,
|
|
target_id,
|
|
reply_to_id=reply_to_id,
|
|
thread_id=thread_id,
|
|
app_id=app_id,
|
|
)
|
|
|
|
async def stream_update(
|
|
self,
|
|
state: CardKitState,
|
|
app_id: str,
|
|
app_secret: str,
|
|
domain: str,
|
|
content: str,
|
|
previous_content: str = "",
|
|
http_timeout_ms: int = 30000,
|
|
) -> str:
|
|
merged = merge_content(previous_content, content)
|
|
if merged == previous_content:
|
|
return previous_content
|
|
|
|
if not should_send_update(merged, previous_content):
|
|
return previous_content
|
|
|
|
try:
|
|
client = get_client(app_id, app_secret, domain, http_timeout_ms)
|
|
await push_card_update(client, state, merged, throttle_ms=self._throttle_ms)
|
|
return merged
|
|
except Exception:
|
|
logger.exception("Card Kit push_card_update failed")
|
|
return previous_content
|
|
|
|
async def stop_stream(
|
|
self,
|
|
state: CardKitState,
|
|
app_id: str,
|
|
app_secret: str,
|
|
domain: str,
|
|
http_timeout_ms: int = 30000,
|
|
) -> None:
|
|
try:
|
|
client = get_client(app_id, app_secret, domain, http_timeout_ms)
|
|
await close_card_stream(client, state)
|
|
except Exception:
|
|
logger.exception("Card Kit close_card_stream failed")
|
|
|
|
async def send_fallback(
|
|
self,
|
|
state: CardKitState,
|
|
content: str,
|
|
app_id: str,
|
|
app_secret: str,
|
|
domain: str,
|
|
http_timeout_ms: int = 30000,
|
|
) -> dict:
|
|
try:
|
|
client = get_client(app_id, app_secret, domain, http_timeout_ms)
|
|
return await send_card_fallback(client, state, content)
|
|
except Exception:
|
|
logger.exception("Card Kit fallback failed")
|
|
return {"success": False, "error": "Internal error"}
|
|
|
|
|
|
class FeishuStreamingRunner:
|
|
def __init__(
|
|
self,
|
|
state: CardKitState,
|
|
app_id: str,
|
|
app_secret: str,
|
|
domain: str,
|
|
http_timeout_ms: int = 30000,
|
|
):
|
|
self._state = state
|
|
self._app_id = app_id
|
|
self._app_secret = app_secret
|
|
self._domain = domain
|
|
self._http_timeout_ms = http_timeout_ms
|
|
self._adapter = FeishuStreamingAdapter()
|
|
self._content = ""
|
|
self._last_pushed = ""
|
|
self._running = False
|
|
self._pending: str | None = None
|
|
self._update_event = asyncio.Event()
|
|
|
|
@property
|
|
def state(self) -> CardKitState:
|
|
return self._state
|
|
|
|
async def feed_delta(self, text: str) -> None:
|
|
self._pending = text
|
|
self._update_event.set()
|
|
|
|
async def feed_end(self) -> None:
|
|
self._running = False
|
|
self._update_event.set()
|
|
|
|
async def freeze_and_reset(self) -> None:
|
|
if self._state.card_id:
|
|
await self._adapter.stop_stream(
|
|
self._state, self._app_id, self._app_secret, self._domain, self._http_timeout_ms
|
|
)
|
|
self._state.card_id = ""
|
|
self._state.message_id = ""
|
|
self._state.version = 0
|
|
self._content = ""
|
|
self._last_pushed = ""
|
|
|
|
async def run(self) -> str:
|
|
self._running = True
|
|
|
|
try:
|
|
while self._running:
|
|
await self._update_event.wait()
|
|
self._update_event.clear()
|
|
|
|
if not self._state.card_id and (self._pending is not None or self._content):
|
|
new_state = await self._adapter.start_stream(
|
|
self._app_id, self._app_secret, self._domain,
|
|
self._state.target_id,
|
|
reply_to_id=self._state.reply_to_id,
|
|
thread_id=self._state.thread_id,
|
|
http_timeout_ms=self._http_timeout_ms,
|
|
)
|
|
if new_state is None:
|
|
continue
|
|
self._state.card_id = new_state.card_id
|
|
self._state.message_id = new_state.message_id
|
|
self._state.version = 0
|
|
|
|
if self._pending is not None:
|
|
self._content += self._pending
|
|
self._pending = None
|
|
for _ in range(3):
|
|
await asyncio.sleep(0)
|
|
if self._pending is not None:
|
|
self._content += self._pending
|
|
self._pending = None
|
|
else:
|
|
break
|
|
|
|
if self._content != self._last_pushed:
|
|
self._last_pushed = await self._adapter.stream_update(
|
|
self._state,
|
|
self._app_id,
|
|
self._app_secret,
|
|
self._domain,
|
|
self._content,
|
|
self._last_pushed,
|
|
self._http_timeout_ms,
|
|
)
|
|
|
|
finally:
|
|
if self._state.card_id:
|
|
await self._adapter.stop_stream(
|
|
self._state,
|
|
self._app_id,
|
|
self._app_secret,
|
|
self._domain,
|
|
self._http_timeout_ms,
|
|
)
|
|
|
|
return self._content
|
|
|
|
|
|
async def stream_reply(
|
|
app_id: str,
|
|
app_secret: str,
|
|
domain: str,
|
|
target_id: str,
|
|
reply_to_id: str | None,
|
|
content_iterator,
|
|
http_timeout_ms: int = 30000,
|
|
) -> dict:
|
|
adapter = FeishuStreamingAdapter()
|
|
|
|
state = await adapter.start_stream(
|
|
app_id, app_secret, domain, target_id,
|
|
reply_to_id=reply_to_id,
|
|
http_timeout_ms=http_timeout_ms,
|
|
)
|
|
|
|
if state is None:
|
|
try:
|
|
client = get_client(app_id, app_secret, domain, http_timeout_ms)
|
|
return await send_card_fallback(
|
|
client,
|
|
CardKitState(target_id=target_id, reply_to_id=reply_to_id),
|
|
content_accumulator(content_iterator),
|
|
)
|
|
except Exception:
|
|
return {"success": False, "error": "Streaming and fallback both failed"}
|
|
|
|
runner = FeishuStreamingRunner(state, app_id, app_secret, domain, http_timeout_ms)
|
|
task = asyncio.create_task(runner.run())
|
|
|
|
try:
|
|
async for chunk in content_iterator:
|
|
await runner.feed_delta(chunk)
|
|
|
|
await runner.feed_end()
|
|
final_content = await task
|
|
except Exception:
|
|
await runner.feed_end()
|
|
await task
|
|
raise
|
|
|
|
return {"success": True, "msg_id": state.message_id, "content": final_content}
|
|
|
|
|
|
async def content_accumulator(content_iterator) -> str:
|
|
parts = []
|
|
async for chunk in content_iterator:
|
|
parts.append(chunk)
|
|
return "".join(parts)
|