Skip to content

Commit e9ae69f

Browse files
authored
fix: prevent SharedPreferences deadlocks (#9649)
* fix: avoid blocking shared preference access * docs: clarify shared preference cache roles * docs: explain shared preference cache purpose * docs: link cache rationale to pull request
1 parent 95181b6 commit e9ae69f

27 files changed

Lines changed: 980 additions & 110 deletions

astrbot/builtin_stars/astrbot/group_chat_context.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -92,7 +92,7 @@ async def get_image_caption(
9292
image_caption_prompt: str,
9393
) -> str:
9494
if not image_caption_provider_id:
95-
provider = self.context.get_using_provider()
95+
provider = await self.context.get_using_provider_async()
9696
else:
9797
provider = self.context.get_provider_by_id(image_caption_provider_id)
9898
if not provider:

astrbot/builtin_stars/astrbot/main.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -227,7 +227,9 @@ async def on_message(self, event: AstrMessageEvent):
227227
logger.error(e)
228228

229229
if need_active:
230-
provider = self.context.get_using_provider(event.unified_msg_origin)
230+
provider = await self.context.get_using_provider_async(
231+
event.unified_msg_origin
232+
)
231233
if not provider:
232234
logger.error("未找到任何 LLM 提供商。请先配置。无法主动回复")
233235
return

astrbot/builtin_stars/builtin_commands/commands/conversation.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -161,7 +161,7 @@ async def reset(self, message: AstrMessageEvent) -> None:
161161
)
162162
return
163163

164-
if not self.context.get_using_provider(umo):
164+
if not await self.context.get_using_provider_async(umo):
165165
message.set_result(
166166
MessageEventResult().message(
167167
"😕 Cannot find any LLM provider. Configure one first."

astrbot/builtin_stars/builtin_commands/commands/provider.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -149,7 +149,7 @@ async def provider(
149149
),
150150
)
151151

152-
provider_using = self.context.get_using_provider(umo=umo)
152+
provider_using = await self.context.get_using_provider_async(umo=umo)
153153
for i, d in enumerate(llm_data):
154154
line = f"{i + 1}. {d['info']}{d['mark']}"
155155
if (
@@ -161,7 +161,7 @@ async def provider(
161161

162162
if tts_data:
163163
parts.append("\n## TTS Providers\n")
164-
tts_using = self.context.get_using_tts_provider(umo=umo)
164+
tts_using = await self.context.get_using_tts_provider_async(umo=umo)
165165
for i, d in enumerate(tts_data):
166166
line = f"{i + 1}. {d['info']}{d['mark']}"
167167
if tts_using and tts_using.meta().id == d["provider"].meta().id:
@@ -170,7 +170,7 @@ async def provider(
170170

171171
if stt_data:
172172
parts.append("\n## STT Providers\n")
173-
stt_using = self.context.get_using_stt_provider(umo=umo)
173+
stt_using = await self.context.get_using_stt_provider_async(umo=umo)
174174
for i, d in enumerate(stt_data):
175175
line = f"{i + 1}. {d['info']}{d['mark']}"
176176
if stt_using and stt_using.meta().id == d["provider"].meta().id:

astrbot/core/astr_main_agent.py

Lines changed: 36 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -227,10 +227,18 @@ def _set_llm_error_message(event: AstrMessageEvent, message: str) -> None:
227227
event.set_extra(LLM_ERROR_MESSAGE_EXTRA_KEY, message)
228228

229229

230-
def _select_provider(
230+
async def _select_provider(
231231
event: AstrMessageEvent, plugin_context: Context
232232
) -> Provider | None:
233-
"""Select chat provider for the event."""
233+
"""Select the chat provider for an event.
234+
235+
Args:
236+
event: Message event that may contain an explicit provider selection.
237+
plugin_context: Plugin context used to resolve configured providers.
238+
239+
Returns:
240+
Selected chat provider, or None if selection fails.
241+
"""
234242
sel_provider = event.get_extra("selected_provider")
235243
if sel_provider and isinstance(sel_provider, str):
236244
provider = plugin_context.get_provider_by_id(sel_provider)
@@ -252,7 +260,9 @@ def _select_provider(
252260
return None
253261
return provider
254262
try:
255-
return plugin_context.get_using_provider(umo=event.unified_msg_origin)
263+
return await plugin_context.get_using_provider_async(
264+
umo=event.unified_msg_origin
265+
)
256266
except ValueError as exc:
257267
logger.error("Error occurred while selecting provider: %s", exc)
258268
_set_llm_error_message(event, f"LLM 请求失败:{exc}")
@@ -916,7 +926,9 @@ async def _process_quote_message(
916926
compress_path = None
917927
prov = plugin_context.get_provider_by_id(img_cap_prov_id)
918928
if prov is None:
919-
prov = plugin_context.get_using_provider(event.unified_msg_origin)
929+
prov = await plugin_context.get_using_provider_async(
930+
event.unified_msg_origin
931+
)
920932

921933
if prov and isinstance(prov, Provider):
922934
path = await image_seg.convert_to_file_path()
@@ -1292,11 +1304,21 @@ def _apply_web_search_citation_prompt(
12921304
req.system_prompt = f"{system_prompt}\n{WEB_SEARCH_CITATION_PROMPT}\n"
12931305

12941306

1295-
def _get_compress_provider(
1307+
async def _get_compress_provider(
12961308
config: MainAgentBuildConfig,
12971309
plugin_context: Context,
12981310
event: AstrMessageEvent | None = None,
12991311
) -> Provider | None:
1312+
"""Resolve the provider used for context compression.
1313+
1314+
Args:
1315+
config: Main agent build configuration.
1316+
plugin_context: Plugin context used to resolve providers.
1317+
event: Optional event used for session-specific fallback selection.
1318+
1319+
Returns:
1320+
Compression provider, or None if compression is disabled or unavailable.
1321+
"""
13001322
if config.context_limit_reached_strategy != "llm_compress":
13011323
return None
13021324
if config.llm_compress_provider_id:
@@ -1310,7 +1332,9 @@ def _get_compress_provider(
13101332
# fallback: use current chat provider for this session
13111333
if event:
13121334
try:
1313-
return plugin_context.get_using_provider(umo=event.unified_msg_origin)
1335+
return await plugin_context.get_using_provider_async(
1336+
umo=event.unified_msg_origin
1337+
)
13141338
except ValueError:
13151339
pass
13161340
return None
@@ -1398,7 +1422,7 @@ async def build_main_agent(
13981422
13991423
If apply_reset is False, will not call reset on the agent runner.
14001424
"""
1401-
provider = provider or _select_provider(event, plugin_context)
1425+
provider = provider or await _select_provider(event, plugin_context)
14021426
if provider is None:
14031427
logger.info("未找到任何对话模型(提供商),跳过 LLM 请求处理。")
14041428
if not event.get_extra(LLM_ERROR_MESSAGE_EXTRA_KEY):
@@ -1699,7 +1723,11 @@ async def build_main_agent(
16991723
streaming=config.streaming_response,
17001724
llm_compress_instruction=config.llm_compress_instruction,
17011725
llm_compress_keep_recent_ratio=config.llm_compress_keep_recent_ratio,
1702-
llm_compress_provider=_get_compress_provider(config, plugin_context, event),
1726+
llm_compress_provider=await _get_compress_provider(
1727+
config,
1728+
plugin_context,
1729+
event,
1730+
),
17031731
truncate_turns=config.dequeue_context_length,
17041732
enforce_max_turns=config.max_context_length,
17051733
tool_schema_mode=config.tool_schema_mode,

astrbot/core/core_lifecycle.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -172,6 +172,8 @@ async def initialize(self) -> None:
172172
LogManager.configure_trace_logger(self.astrbot_config)
173173

174174
await self.db.initialize()
175+
if sp.db_helper is self.db:
176+
await sp.initialize()
175177

176178
await html_renderer.initialize()
177179

@@ -404,6 +406,8 @@ async def stop(self) -> None:
404406
await self.provider_manager.terminate()
405407
await self.platform_manager.terminate()
406408
await self.kb_manager.terminate()
409+
if sp.db_helper is self.db:
410+
await sp.close()
407411
self.dashboard_shutdown_event.set()
408412

409413
# 再次遍历curr_tasks等待每个任务真正结束
@@ -427,6 +431,8 @@ async def restart(self) -> None:
427431
await self.provider_manager.terminate()
428432
await self.platform_manager.terminate()
429433
await self.kb_manager.terminate()
434+
if sp.db_helper is self.db:
435+
await sp.close()
430436
self.dashboard_shutdown_event.set()
431437
threading.Thread(
432438
target=restart_process,

astrbot/core/db/__init__.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -577,11 +577,11 @@ async def get_preference(self, scope: str, scope_id: str, key: str) -> Preferenc
577577
@abc.abstractmethod
578578
async def get_preferences(
579579
self,
580-
scope: str,
580+
scope: str | None = None,
581581
scope_id: str | None = None,
582582
key: str | None = None,
583583
) -> list[Preference]:
584-
"""Get all preferences for a specific scope ID or key."""
584+
"""Get preferences, optionally filtered by scope, scope ID, or key."""
585585
...
586586

587587
@abc.abstractmethod

astrbot/core/db/sqlite.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1379,11 +1379,13 @@ async def get_preference(self, scope, scope_id, key):
13791379
result = await session.execute(query)
13801380
return result.scalar_one_or_none()
13811381

1382-
async def get_preferences(self, scope, scope_id=None, key=None):
1383-
"""Get all preferences for a specific scope ID or key."""
1382+
async def get_preferences(self, scope=None, scope_id=None, key=None):
1383+
"""Get preferences, optionally filtered by scope, scope ID, or key."""
13841384
async with self.get_db() as session:
13851385
session: AsyncSession
1386-
query = select(Preference).where(Preference.scope == scope)
1386+
query = select(Preference)
1387+
if scope is not None:
1388+
query = query.where(Preference.scope == scope)
13871389
if scope_id is not None:
13881390
query = query.where(Preference.scope_id == scope_id)
13891391
if key is not None:

astrbot/core/pipeline/preprocess_stage/stage.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -169,7 +169,9 @@ async def process(
169169
if self.stt_settings.get("enable", False):
170170
# TODO: 独立
171171
ctx = self.plugin_manager.context
172-
stt_provider = ctx.get_using_stt_provider(event.unified_msg_origin)
172+
stt_provider = await ctx.get_using_stt_provider_async(
173+
event.unified_msg_origin
174+
)
173175
if not stt_provider:
174176
logger.warning(
175177
f"Session {event.unified_msg_origin} has no speech-to-text "

astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -298,10 +298,8 @@ async def process(
298298
)
299299

300300
# 获取 TTS Provider
301-
tts_provider = (
302-
self.ctx.plugin_manager.context.get_using_tts_provider(
303-
event.unified_msg_origin
304-
)
301+
tts_provider = await self.ctx.plugin_manager.context.get_using_tts_provider_async(
302+
event.unified_msg_origin
305303
)
306304

307305
if not tts_provider:

0 commit comments

Comments
 (0)