mirror of
https://github.com/HKUDS/OpenSpace.git
synced 2026-10-08 03:07:51 +00:00
feat: add communication adapters and agent refactoring (#74)
This commit is contained in:
parent
4791133e11
commit
f01d408ac4
27 changed files with 4791 additions and 338 deletions
5
.gitignore
vendored
5
.gitignore
vendored
|
|
@ -28,9 +28,12 @@ build/
|
|||
.env
|
||||
!.env.example
|
||||
|
||||
# MCP files
|
||||
# MCP config files
|
||||
openspace/config/config_mcp.json
|
||||
|
||||
# Communication config files
|
||||
openspace/config/config_communication.json
|
||||
|
||||
# Logs
|
||||
logs/
|
||||
|
||||
|
|
|
|||
|
|
@ -38,6 +38,13 @@ OPENROUTER_API_KEY=
|
|||
# OPENSPACE_LLM_API_KEY=sk-xxx
|
||||
# OPENSPACE_LLM_API_BASE=https://openrouter.ai/api/v1
|
||||
|
||||
# --- Option C: Local Ollama ---
|
||||
# For ollama/* models, set OPENSPACE_MODEL and the local Ollama endpoint.
|
||||
#
|
||||
# OPENSPACE_MODEL=ollama/qwen3-coder:30b
|
||||
# OLLAMA_API_BASE=http://127.0.0.1:11434
|
||||
# OLLAMA_API_KEY=ollama
|
||||
|
||||
# ── OpenSpace Cloud (optional) ──────────────────────────────
|
||||
# Register at https://open-space.cloud to get your key.
|
||||
# Enables cloud skill search & upload; local features work without it.
|
||||
|
|
|
|||
|
|
@ -158,6 +158,43 @@ def _create_argument_parser() -> argparse.ArgumentParser:
|
|||
'--config', '-c', type=str,
|
||||
help='MCP configuration file path'
|
||||
)
|
||||
|
||||
communication_parser = subparsers.add_parser(
|
||||
'communication',
|
||||
help='Run the communication gateway'
|
||||
)
|
||||
communication_parser.add_argument(
|
||||
'--config',
|
||||
type=str,
|
||||
dest='communication_config',
|
||||
help='Communication configuration file path'
|
||||
)
|
||||
communication_subparsers = communication_parser.add_subparsers(
|
||||
dest='communication_command',
|
||||
help='Communication gateway commands'
|
||||
)
|
||||
communication_run_parser = communication_subparsers.add_parser(
|
||||
'run',
|
||||
help='Start the communication gateway'
|
||||
)
|
||||
communication_run_parser.add_argument(
|
||||
'--config',
|
||||
type=str,
|
||||
dest='communication_config',
|
||||
help='Communication configuration file path'
|
||||
)
|
||||
communication_health_parser = communication_subparsers.add_parser(
|
||||
'health',
|
||||
help='Check the communication gateway health endpoint'
|
||||
)
|
||||
communication_health_parser.add_argument(
|
||||
'--config',
|
||||
type=str,
|
||||
dest='communication_config',
|
||||
help='Communication configuration file path'
|
||||
)
|
||||
communication_health_parser.add_argument('--host', type=str, default=None)
|
||||
communication_health_parser.add_argument('--port', type=int, default=None)
|
||||
|
||||
# Basic arguments (for run mode)
|
||||
parser.add_argument('--config', '-c', type=str, help='Configuration file path (JSON format)')
|
||||
|
|
@ -444,6 +481,20 @@ async def main():
|
|||
if args.command == 'refresh-cache':
|
||||
await refresh_mcp_cache(args.config)
|
||||
return 0
|
||||
if args.command == 'communication':
|
||||
from openspace.communication.gateway import main as communication_main
|
||||
|
||||
communication_argv = []
|
||||
if args.communication_config:
|
||||
communication_argv.extend(['--config', args.communication_config])
|
||||
if args.communication_command:
|
||||
communication_argv.append(args.communication_command)
|
||||
if args.communication_command == 'health':
|
||||
if args.host:
|
||||
communication_argv.extend(['--host', args.host])
|
||||
if args.port is not None:
|
||||
communication_argv.extend(['--port', str(args.port)])
|
||||
return await communication_main(communication_argv)
|
||||
|
||||
# Load configuration
|
||||
config = _load_config(args)
|
||||
|
|
|
|||
|
|
@ -5,8 +5,15 @@ import json
|
|||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
from openspace.agents.base import BaseAgent
|
||||
from openspace.agents.message_utils import (
|
||||
ITERATION_GUIDANCE_PREFIX,
|
||||
build_channel_context_message,
|
||||
cap_message_content,
|
||||
normalize_external_history,
|
||||
truncate_messages,
|
||||
)
|
||||
from openspace.agents.visual_analyzer import VisualAnalyzer
|
||||
from openspace.grounding.core.types import BackendType, ToolResult
|
||||
from openspace.platforms.screenshot import ScreenshotClient
|
||||
from openspace.prompts import GroundingAgentPrompts
|
||||
from openspace.utils.logging import Logger
|
||||
|
||||
|
|
@ -20,6 +27,7 @@ logger = Logger.get_logger(__name__)
|
|||
|
||||
|
||||
class GroundingAgent(BaseAgent):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
name: str = "GroundingAgent",
|
||||
|
|
@ -58,9 +66,12 @@ class GroundingAgent(BaseAgent):
|
|||
|
||||
self._system_prompt = system_prompt or self._default_system_prompt()
|
||||
self._max_iterations = max_iterations
|
||||
self._visual_analysis_timeout = visual_analysis_timeout
|
||||
self._tool_retrieval_llm = tool_retrieval_llm
|
||||
self._visual_analysis_model = visual_analysis_model
|
||||
self._visual_analyzer = VisualAnalyzer(
|
||||
llm_client=llm_client,
|
||||
visual_analysis_model=visual_analysis_model,
|
||||
visual_analysis_timeout=visual_analysis_timeout,
|
||||
)
|
||||
|
||||
# Skill context injection (set externally before process())
|
||||
self._skill_context: Optional[str] = None
|
||||
|
|
@ -75,7 +86,7 @@ class GroundingAgent(BaseAgent):
|
|||
logger.info(f"Grounding Agent initialized: {name}")
|
||||
logger.info(f"Backend scope: {self._backend_scope}")
|
||||
logger.info(f"Max iterations: {self._max_iterations}")
|
||||
logger.info(f"Visual analysis timeout: {self._visual_analysis_timeout}s")
|
||||
logger.info(f"Visual analysis timeout: {visual_analysis_timeout}s")
|
||||
if tool_retrieval_llm:
|
||||
logger.info(f"Tool retrieval model: {tool_retrieval_llm.model}")
|
||||
if visual_analysis_model:
|
||||
|
|
@ -119,92 +130,6 @@ class GroundingAgent(BaseAgent):
|
|||
count = len(registry.list_skills())
|
||||
logger.info(f"Skill registry attached ({count} skill(s) available for mid-iteration retrieval)")
|
||||
|
||||
_MAX_SINGLE_CONTENT_CHARS = 30_000
|
||||
_ITERATION_GUIDANCE_PREFIX = "[INTERNAL ORCHESTRATION NOTE]"
|
||||
|
||||
@classmethod
|
||||
def _cap_message_content(cls, messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
||||
"""Truncate oversized individual message contents in-place.
|
||||
|
||||
Targets tool-result messages and assistant messages that can
|
||||
carry enormous file contents (read_file on large CSVs/scripts).
|
||||
System messages and the first user instruction are never touched.
|
||||
"""
|
||||
cap = cls._MAX_SINGLE_CONTENT_CHARS
|
||||
trimmed = 0
|
||||
for msg in messages:
|
||||
content = msg.get("content")
|
||||
if not isinstance(content, str) or len(content) <= cap:
|
||||
continue
|
||||
if msg.get("role") == "system":
|
||||
continue
|
||||
original_len = len(content)
|
||||
msg["content"] = (
|
||||
content[: cap // 2]
|
||||
+ f"\n\n... [truncated {original_len - cap:,} chars] ...\n\n"
|
||||
+ content[-(cap // 2):]
|
||||
)
|
||||
trimmed += 1
|
||||
if trimmed:
|
||||
logger.info(f"Capped {trimmed} oversized message(s) to {cap:,} chars each")
|
||||
return messages
|
||||
|
||||
def _truncate_messages(
|
||||
self,
|
||||
messages: List[Dict[str, Any]],
|
||||
keep_recent: int = 8,
|
||||
max_tokens_estimate: int = 120000
|
||||
) -> List[Dict[str, Any]]:
|
||||
# First: cap any single oversized message to prevent one huge
|
||||
# tool-result from dominating the context window.
|
||||
messages = self._cap_message_content(messages)
|
||||
|
||||
if len(messages) <= keep_recent + 2: # +2 for system and initial user
|
||||
return messages
|
||||
|
||||
total_text = json.dumps(messages, ensure_ascii=False)
|
||||
estimated_tokens = len(total_text) // 4
|
||||
|
||||
if estimated_tokens < max_tokens_estimate:
|
||||
return messages
|
||||
|
||||
logger.info(f"Truncating message history: {len(messages)} messages, "
|
||||
f"~{estimated_tokens:,} tokens -> keeping recent {keep_recent} rounds")
|
||||
|
||||
system_messages = []
|
||||
user_instruction = None
|
||||
conversation_messages = []
|
||||
|
||||
for msg in messages:
|
||||
role = msg.get("role")
|
||||
if role == "system":
|
||||
system_messages.append(msg)
|
||||
elif role == "user" and user_instruction is None:
|
||||
user_instruction = msg
|
||||
else:
|
||||
conversation_messages.append(msg)
|
||||
|
||||
recent_messages = conversation_messages[-(keep_recent * 2):] if conversation_messages else []
|
||||
|
||||
truncated = system_messages.copy()
|
||||
dropped = len(conversation_messages) - len(recent_messages)
|
||||
if dropped > 0:
|
||||
truncated.append({
|
||||
"role": "system",
|
||||
"content": (
|
||||
f"{self._ITERATION_GUIDANCE_PREFIX} {dropped} earlier messages were "
|
||||
"truncated to save context. The original task instruction is preserved below."
|
||||
),
|
||||
})
|
||||
if user_instruction:
|
||||
truncated.append(user_instruction)
|
||||
truncated.extend(recent_messages)
|
||||
|
||||
logger.info(f"After truncation: {len(truncated)} messages, "
|
||||
f"~{len(json.dumps(truncated, ensure_ascii=False))//4:,} tokens (estimated)")
|
||||
|
||||
return truncated
|
||||
|
||||
async def process(self, context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Process a task execution request with multi-round iteration control.
|
||||
|
|
@ -286,6 +211,14 @@ class GroundingAgent(BaseAgent):
|
|||
tools=tools,
|
||||
)
|
||||
|
||||
async def _va_callback(
|
||||
result: ToolResult, tool_name: str, tool_call: Dict, backend: str
|
||||
) -> ToolResult:
|
||||
return await self._visual_analyzer.analyze_tool_result(
|
||||
result, tool_name, tool_call, backend,
|
||||
task_description=instruction,
|
||||
)
|
||||
|
||||
try:
|
||||
while current_iteration < max_iterations:
|
||||
current_iteration += 1
|
||||
|
|
@ -305,15 +238,15 @@ class GroundingAgent(BaseAgent):
|
|||
# Cap oversized individual messages every iteration to prevent
|
||||
# a single huge tool result from ballooning all subsequent calls.
|
||||
if current_iteration >= 2:
|
||||
messages = self._cap_message_content(messages)
|
||||
messages = cap_message_content(messages)
|
||||
|
||||
# Truncate message history to prevent context length issues
|
||||
# Start truncating after 5 iterations to keep context manageable
|
||||
if current_iteration >= 5:
|
||||
messages = self._truncate_messages(
|
||||
messages,
|
||||
messages = truncate_messages(
|
||||
messages,
|
||||
keep_recent=8,
|
||||
max_tokens_estimate=120000
|
||||
max_tokens_estimate=120000,
|
||||
)
|
||||
|
||||
messages_input_snapshot = copy.deepcopy(messages)
|
||||
|
|
@ -335,7 +268,7 @@ class GroundingAgent(BaseAgent):
|
|||
tools=tools if context.get("auto_execute", True) else None,
|
||||
execute_tools=context.get("auto_execute", True),
|
||||
summary_prompt=None, # Disabled
|
||||
tool_result_callback=self._visual_analysis_callback
|
||||
tool_result_callback=_va_callback,
|
||||
)
|
||||
|
||||
# Update messages with LLM response
|
||||
|
|
@ -427,7 +360,7 @@ class GroundingAgent(BaseAgent):
|
|||
msg for msg in messages
|
||||
if not (
|
||||
isinstance(msg.get("content"), str)
|
||||
and msg.get("content", "").startswith(self._ITERATION_GUIDANCE_PREFIX)
|
||||
and msg.get("content", "").startswith(ITERATION_GUIDANCE_PREFIX)
|
||||
)
|
||||
]
|
||||
|
||||
|
|
@ -435,7 +368,7 @@ class GroundingAgent(BaseAgent):
|
|||
# so runtime guidance is sent as an internal user note.
|
||||
guidance_msg = {
|
||||
"role": "user",
|
||||
"content": f"{self._ITERATION_GUIDANCE_PREFIX}\n"
|
||||
"content": f"{ITERATION_GUIDANCE_PREFIX}\n"
|
||||
f"Iteration {current_iteration} complete. "
|
||||
f"Check if task is finished - if yes, output {GroundingAgentPrompts.TASK_COMPLETE}. "
|
||||
f"If not, continue with next action."
|
||||
|
|
@ -532,7 +465,16 @@ class GroundingAgent(BaseAgent):
|
|||
"role": "system",
|
||||
"content": artifact_msg
|
||||
})
|
||||
|
||||
|
||||
channel_context_msg = build_channel_context_message(
|
||||
context.get("channel_context")
|
||||
)
|
||||
if channel_context_msg:
|
||||
messages.append({
|
||||
"role": "system",
|
||||
"content": channel_context_msg,
|
||||
})
|
||||
|
||||
# Skill injection — only active (selected) skills, full content
|
||||
if self._skill_context:
|
||||
messages.append({
|
||||
|
|
@ -540,10 +482,20 @@ class GroundingAgent(BaseAgent):
|
|||
"content": self._skill_context
|
||||
})
|
||||
logger.info(f"Injected active skill context ({len(self._active_skill_ids)} skill(s))")
|
||||
|
||||
|
||||
external_history = normalize_external_history(
|
||||
context.get("conversation_history")
|
||||
)
|
||||
if external_history:
|
||||
messages.extend(external_history)
|
||||
logger.info(
|
||||
"Injected %d external conversation message(s)",
|
||||
len(external_history),
|
||||
)
|
||||
|
||||
# User instruction
|
||||
messages.append({"role": "user", "content": instruction})
|
||||
|
||||
|
||||
return messages
|
||||
|
||||
async def _get_available_tools(self, task_description: Optional[str]) -> List:
|
||||
|
|
@ -629,230 +581,6 @@ class GroundingAgent(BaseAgent):
|
|||
)
|
||||
return all_tools
|
||||
|
||||
async def _visual_analysis_callback(
|
||||
self,
|
||||
result: ToolResult,
|
||||
tool_name: str,
|
||||
tool_call: Dict,
|
||||
backend: str
|
||||
) -> ToolResult:
|
||||
"""
|
||||
Callback for LLMClient to handle visual analysis after tool execution.
|
||||
"""
|
||||
# 1. Check if LLM requested to skip visual analysis
|
||||
skip_visual_analysis = False
|
||||
try:
|
||||
arguments = tool_call.function.arguments
|
||||
if isinstance(arguments, str):
|
||||
args = json.loads(arguments.strip() or "{}")
|
||||
else:
|
||||
args = arguments
|
||||
|
||||
if isinstance(args, dict) and args.get("skip_visual_analysis"):
|
||||
skip_visual_analysis = True
|
||||
logger.info(f"Visual analysis skipped for {tool_name} (meta-parameter set by LLM)")
|
||||
except Exception as e:
|
||||
logger.debug(f"Could not parse tool arguments: {e}")
|
||||
|
||||
# 2. If skip requested, return original result
|
||||
if skip_visual_analysis:
|
||||
return result
|
||||
|
||||
# 3. Check if this backend needs visual analysis
|
||||
if backend != "gui":
|
||||
return result
|
||||
|
||||
# 4. Check if tool has visual data
|
||||
metadata = getattr(result, 'metadata', None)
|
||||
has_screenshots = metadata and (metadata.get("screenshot") or metadata.get("screenshots"))
|
||||
|
||||
# 5. If no visual data, try to capture a screenshot
|
||||
if not has_screenshots:
|
||||
try:
|
||||
logger.info(f"No visual data from {tool_name}, capturing screenshot...")
|
||||
screenshot_client = ScreenshotClient()
|
||||
screenshot_bytes = await screenshot_client.capture()
|
||||
|
||||
if screenshot_bytes:
|
||||
# Add screenshot to result metadata
|
||||
if metadata is None:
|
||||
result.metadata = {}
|
||||
metadata = result.metadata
|
||||
metadata["screenshot"] = screenshot_bytes
|
||||
has_screenshots = True
|
||||
logger.info(f"Screenshot captured for visual analysis")
|
||||
else:
|
||||
logger.warning("Failed to capture screenshot")
|
||||
except Exception as e:
|
||||
logger.warning(f"Error capturing screenshot: {e}")
|
||||
|
||||
# 6. If still no screenshots, return original result
|
||||
if not has_screenshots:
|
||||
logger.debug(f"No visual data available for {tool_name}")
|
||||
return result
|
||||
|
||||
# 7. Perform visual analysis
|
||||
return await self._enhance_result_with_visual_context(result, tool_name)
|
||||
|
||||
async def _enhance_result_with_visual_context(
|
||||
self,
|
||||
result: ToolResult,
|
||||
tool_name: str
|
||||
) -> ToolResult:
|
||||
"""
|
||||
Enhance tool result with visual analysis for grounding agent workflows.
|
||||
"""
|
||||
import asyncio
|
||||
import base64
|
||||
import litellm
|
||||
|
||||
try:
|
||||
metadata = getattr(result, 'metadata', None)
|
||||
if not metadata:
|
||||
return result
|
||||
|
||||
# Collect all screenshots
|
||||
screenshots_bytes = []
|
||||
|
||||
# Check for multiple screenshots first
|
||||
if metadata.get("screenshots"):
|
||||
screenshots_list = metadata["screenshots"]
|
||||
if isinstance(screenshots_list, list):
|
||||
screenshots_bytes = [s for s in screenshots_list if s]
|
||||
# Fall back to single screenshot
|
||||
elif metadata.get("screenshot"):
|
||||
screenshots_bytes = [metadata["screenshot"]]
|
||||
|
||||
if not screenshots_bytes:
|
||||
return result
|
||||
|
||||
# Select key screenshots if there are too many
|
||||
selected_screenshots = self._select_key_screenshots(screenshots_bytes, max_count=3)
|
||||
|
||||
# Convert to base64
|
||||
visual_b64_list = []
|
||||
for visual_data in selected_screenshots:
|
||||
if isinstance(visual_data, bytes):
|
||||
visual_b64_list.append(base64.b64encode(visual_data).decode('utf-8'))
|
||||
else:
|
||||
visual_b64_list.append(visual_data) # Already base64
|
||||
|
||||
# Build prompt based on number of screenshots
|
||||
num_screenshots = len(visual_b64_list)
|
||||
|
||||
prompt = GroundingAgentPrompts.visual_analysis(
|
||||
tool_name=tool_name,
|
||||
num_screenshots=num_screenshots,
|
||||
task_description=getattr(self, '_current_instruction', '')
|
||||
)
|
||||
|
||||
# Build content with text prompt + all images
|
||||
content = [{"type": "text", "text": prompt}]
|
||||
for visual_b64 in visual_b64_list:
|
||||
content.append({
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": f"data:image/png;base64,{visual_b64}"
|
||||
}
|
||||
})
|
||||
|
||||
# Resolve visual-model credentials independently when the visual
|
||||
# model differs from the main reasoning model.
|
||||
visual_model = self._visual_analysis_model or (self._llm_client.model if self._llm_client else "openrouter/anthropic/claude-sonnet-4.5")
|
||||
_llm_extra = {}
|
||||
if self._llm_client and visual_model == self._llm_client.model:
|
||||
_llm_extra = getattr(self._llm_client, 'litellm_kwargs', {}) or {}
|
||||
elif self._visual_analysis_model:
|
||||
try:
|
||||
from openspace.host_detection import build_llm_kwargs
|
||||
visual_model, _llm_extra = build_llm_kwargs(visual_model)
|
||||
except Exception as e:
|
||||
logger.debug(f"Failed to resolve dedicated visual model credentials: {e}")
|
||||
_llm_extra = {}
|
||||
response = await asyncio.wait_for(
|
||||
litellm.acompletion(
|
||||
model=visual_model,
|
||||
messages=[{
|
||||
"role": "user",
|
||||
"content": content
|
||||
}],
|
||||
timeout=self._visual_analysis_timeout,
|
||||
**_llm_extra,
|
||||
),
|
||||
timeout=self._visual_analysis_timeout + 5
|
||||
)
|
||||
|
||||
analysis = response.choices[0].message.content.strip()
|
||||
|
||||
# Inject visual analysis into content
|
||||
original_content = result.content or "(no text output)"
|
||||
enhanced_content = f"{original_content}\n\n**Visual content**: {analysis}"
|
||||
|
||||
# Create enhanced result
|
||||
enhanced_result = ToolResult(
|
||||
status=result.status,
|
||||
content=enhanced_content,
|
||||
error=result.error,
|
||||
metadata={**metadata, "visual_analyzed": True, "visual_analysis": analysis},
|
||||
execution_time=result.execution_time
|
||||
)
|
||||
|
||||
logger.info(f"Enhanced {tool_name} result with visual analysis ({num_screenshots} screenshot(s))")
|
||||
return enhanced_result
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning(f"Visual analysis timed out for {tool_name}, returning original result")
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to analyze visual content for {tool_name}: {e}")
|
||||
return result
|
||||
|
||||
def _select_key_screenshots(
|
||||
self,
|
||||
screenshots: List[bytes],
|
||||
max_count: int = 3
|
||||
) -> List[bytes]:
|
||||
"""
|
||||
Select key screenshots if there are too many.
|
||||
"""
|
||||
if len(screenshots) <= max_count:
|
||||
return screenshots
|
||||
|
||||
selected_indices = set()
|
||||
|
||||
# Always include last (final state)
|
||||
selected_indices.add(len(screenshots) - 1)
|
||||
|
||||
# If room, include first (initial state)
|
||||
if max_count >= 2:
|
||||
selected_indices.add(0)
|
||||
|
||||
# Fill remaining slots with evenly spaced middle screenshots
|
||||
remaining_slots = max_count - len(selected_indices)
|
||||
if remaining_slots > 0:
|
||||
# Calculate spacing
|
||||
available_indices = [
|
||||
i for i in range(1, len(screenshots) - 1)
|
||||
if i not in selected_indices
|
||||
]
|
||||
|
||||
if available_indices:
|
||||
step = max(1, len(available_indices) // (remaining_slots + 1))
|
||||
for i in range(remaining_slots):
|
||||
idx = min((i + 1) * step, len(available_indices) - 1)
|
||||
if idx < len(available_indices):
|
||||
selected_indices.add(available_indices[idx])
|
||||
|
||||
# Return screenshots in original order
|
||||
selected = [screenshots[i] for i in sorted(selected_indices)]
|
||||
|
||||
logger.debug(
|
||||
f"Selected {len(selected)} screenshots at indices {sorted(selected_indices)} "
|
||||
f"from total of {len(screenshots)}"
|
||||
)
|
||||
|
||||
return selected
|
||||
|
||||
def _get_workspace_path(self, context: Dict[str, Any]) -> Optional[str]:
|
||||
"""
|
||||
Get workspace directory path from context.
|
||||
|
|
|
|||
227
openspace/agents/message_utils.py
Normal file
227
openspace/agents/message_utils.py
Normal file
|
|
@ -0,0 +1,227 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional, Set
|
||||
|
||||
from openspace.prompts import GroundingAgentPrompts
|
||||
from openspace.utils.logging import Logger
|
||||
|
||||
logger = Logger.get_logger(__name__)
|
||||
|
||||
SUPPORTED_EXTERNAL_HISTORY_ROLES: Set[str] = {"user", "assistant"}
|
||||
MAX_SINGLE_CONTENT_CHARS: int = 30_000
|
||||
ITERATION_GUIDANCE_PREFIX: str = "[INTERNAL ORCHESTRATION NOTE]"
|
||||
|
||||
|
||||
def cap_message_content(
|
||||
messages: List[Dict[str, Any]],
|
||||
max_chars: int = MAX_SINGLE_CONTENT_CHARS,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Truncate oversized individual message contents in-place.
|
||||
|
||||
Targets tool-result messages and assistant messages that can
|
||||
carry enormous file contents (read_file on large CSVs/scripts).
|
||||
System messages and the first user instruction are never touched.
|
||||
"""
|
||||
trimmed = 0
|
||||
for msg in messages:
|
||||
content = msg.get("content")
|
||||
if not isinstance(content, str) or len(content) <= max_chars:
|
||||
continue
|
||||
if msg.get("role") == "system":
|
||||
continue
|
||||
original_len = len(content)
|
||||
msg["content"] = (
|
||||
content[: max_chars // 2]
|
||||
+ f"\n\n... [truncated {original_len - max_chars:,} chars] ...\n\n"
|
||||
+ content[-(max_chars // 2) :]
|
||||
)
|
||||
trimmed += 1
|
||||
if trimmed:
|
||||
logger.info(f"Capped {trimmed} oversized message(s) to {max_chars:,} chars each")
|
||||
return messages
|
||||
|
||||
|
||||
def truncate_messages(
|
||||
messages: List[Dict[str, Any]],
|
||||
keep_recent: int = 8,
|
||||
max_tokens_estimate: int = 120_000,
|
||||
guidance_prefix: str = ITERATION_GUIDANCE_PREFIX,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Truncate conversation history to fit within token budget.
|
||||
|
||||
Preserves system messages and the first user instruction while
|
||||
keeping only the most recent conversation turns.
|
||||
"""
|
||||
messages = cap_message_content(messages)
|
||||
|
||||
if len(messages) <= keep_recent + 2: # +2 for system and initial user
|
||||
return messages
|
||||
|
||||
total_text = json.dumps(messages, ensure_ascii=False)
|
||||
estimated_tokens = len(total_text) // 4
|
||||
|
||||
if estimated_tokens < max_tokens_estimate:
|
||||
return messages
|
||||
|
||||
logger.info(
|
||||
f"Truncating message history: {len(messages)} messages, "
|
||||
f"~{estimated_tokens:,} tokens -> keeping recent {keep_recent} rounds"
|
||||
)
|
||||
|
||||
system_messages: List[Dict[str, Any]] = []
|
||||
user_instruction: Optional[Dict[str, Any]] = None
|
||||
conversation_messages: List[Dict[str, Any]] = []
|
||||
|
||||
for msg in messages:
|
||||
role = msg.get("role")
|
||||
if role == "system":
|
||||
system_messages.append(msg)
|
||||
elif role == "user" and user_instruction is None:
|
||||
user_instruction = msg
|
||||
else:
|
||||
conversation_messages.append(msg)
|
||||
|
||||
recent_messages = (
|
||||
conversation_messages[-(keep_recent * 2) :] if conversation_messages else []
|
||||
)
|
||||
|
||||
truncated = system_messages.copy()
|
||||
dropped = len(conversation_messages) - len(recent_messages)
|
||||
if dropped > 0:
|
||||
truncated.append(
|
||||
{
|
||||
"role": "system",
|
||||
"content": (
|
||||
f"{guidance_prefix} {dropped} earlier messages were "
|
||||
"truncated to save context. The original task instruction "
|
||||
"is preserved below."
|
||||
),
|
||||
}
|
||||
)
|
||||
if user_instruction:
|
||||
truncated.append(user_instruction)
|
||||
truncated.extend(recent_messages)
|
||||
|
||||
logger.info(
|
||||
f"After truncation: {len(truncated)} messages, "
|
||||
f"~{len(json.dumps(truncated, ensure_ascii=False)) // 4:,} tokens (estimated)"
|
||||
)
|
||||
|
||||
return truncated
|
||||
|
||||
|
||||
def normalize_external_history(
|
||||
conversation_history: Any,
|
||||
supported_roles: Set[str] = SUPPORTED_EXTERNAL_HISTORY_ROLES,
|
||||
) -> List[Dict[str, str]]:
|
||||
"""Normalize external conversation history into ``{role, content}`` dicts."""
|
||||
if not isinstance(conversation_history, list):
|
||||
return []
|
||||
|
||||
normalized: List[Dict[str, str]] = []
|
||||
for entry in conversation_history:
|
||||
if not isinstance(entry, dict):
|
||||
continue
|
||||
role = str(entry.get("role", "")).strip().lower()
|
||||
if role not in supported_roles:
|
||||
continue
|
||||
|
||||
content = entry.get("content")
|
||||
if isinstance(content, list):
|
||||
parts: List[str] = []
|
||||
for item in content:
|
||||
if isinstance(item, dict):
|
||||
text = item.get("text")
|
||||
if isinstance(text, str) and text.strip():
|
||||
parts.append(text.strip())
|
||||
elif isinstance(item, str) and item.strip():
|
||||
parts.append(item.strip())
|
||||
content = "\n".join(parts).strip()
|
||||
elif content is not None:
|
||||
content = str(content).strip()
|
||||
|
||||
if not content:
|
||||
continue
|
||||
|
||||
normalized.append({"role": role, "content": content})
|
||||
|
||||
return normalized
|
||||
|
||||
|
||||
def build_channel_context_message(channel_context: Any) -> Optional[str]:
|
||||
"""Build a system message describing the communication channel context."""
|
||||
if not isinstance(channel_context, dict):
|
||||
return None
|
||||
|
||||
lines = [
|
||||
"## Channel Context",
|
||||
]
|
||||
|
||||
platform = str(channel_context.get("platform", "")).strip()
|
||||
chat_type = str(channel_context.get("chat_type", "")).strip()
|
||||
chat_id = str(channel_context.get("chat_id", "")).strip()
|
||||
chat_name = str(channel_context.get("chat_name", "")).strip()
|
||||
thread_id = str(channel_context.get("thread_id", "")).strip()
|
||||
user_name = str(channel_context.get("user_name", "")).strip()
|
||||
user_id = str(channel_context.get("user_id", "")).strip()
|
||||
session_key = str(channel_context.get("session_key", "")).strip()
|
||||
message_id = str(channel_context.get("message_id", "")).strip()
|
||||
reply_to_message_id = str(channel_context.get("reply_to_message_id", "")).strip()
|
||||
reply_to_text = str(channel_context.get("reply_to_text", "")).strip()
|
||||
|
||||
if platform:
|
||||
lines.append(f"- Platform: {platform}")
|
||||
if chat_type:
|
||||
lines.append(f"- Chat type: {chat_type}")
|
||||
if chat_id:
|
||||
lines.append(f"- Chat ID: {chat_id}")
|
||||
if chat_name:
|
||||
lines.append(f"- Chat name: {chat_name}")
|
||||
if thread_id:
|
||||
lines.append(f"- Thread ID: {thread_id}")
|
||||
if user_name:
|
||||
lines.append(f"- User: {user_name}")
|
||||
elif user_id:
|
||||
lines.append(f"- User ID: {user_id}")
|
||||
if session_key:
|
||||
lines.append(f"- Session key: {session_key}")
|
||||
if message_id:
|
||||
lines.append(f"- Message ID: {message_id}")
|
||||
if reply_to_message_id:
|
||||
lines.append(f"- Reply-to message ID: {reply_to_message_id}")
|
||||
if reply_to_text:
|
||||
lines.append(f"- Reply context: {reply_to_text[:500]}")
|
||||
|
||||
lines.extend(
|
||||
[
|
||||
"",
|
||||
"## Chat Reply Policy",
|
||||
"- If the user is making simple conversation, answer directly in natural language.",
|
||||
"- Do not call tools for greetings, acknowledgements, thanks, or brief "
|
||||
"clarifications that can be answered from the current context.",
|
||||
f"- When you reply directly without tools, include "
|
||||
f"`{GroundingAgentPrompts.TASK_COMPLETE}` at the end of your response.",
|
||||
]
|
||||
)
|
||||
|
||||
attachments = channel_context.get("attachments")
|
||||
if isinstance(attachments, list) and attachments:
|
||||
lines.append("- Attachments:")
|
||||
for attachment in attachments:
|
||||
if not isinstance(attachment, dict):
|
||||
continue
|
||||
path = str(attachment.get("path", "")).strip()
|
||||
if not path:
|
||||
continue
|
||||
kind = str(attachment.get("kind", "file")).strip() or "file"
|
||||
name = str(attachment.get("name", "")).strip()
|
||||
label = f"{kind}: {path}"
|
||||
if name:
|
||||
label += f" ({name})"
|
||||
lines.append(f" - {label}")
|
||||
|
||||
if len(lines) == 1:
|
||||
return None
|
||||
|
||||
return "\n".join(lines)
|
||||
250
openspace/agents/visual_analyzer.py
Normal file
250
openspace/agents/visual_analyzer.py
Normal file
|
|
@ -0,0 +1,250 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
from openspace.grounding.core.types import ToolResult
|
||||
from openspace.platforms.screenshot import ScreenshotClient
|
||||
from openspace.prompts import GroundingAgentPrompts
|
||||
from openspace.utils.logging import Logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from openspace.llm import LLMClient
|
||||
|
||||
logger = Logger.get_logger(__name__)
|
||||
|
||||
|
||||
class VisualAnalyzer:
|
||||
"""Handles screenshot capture and LLM-based visual analysis of tool results."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
llm_client: Optional[LLMClient] = None,
|
||||
visual_analysis_model: Optional[str] = None,
|
||||
visual_analysis_timeout: float = 30.0,
|
||||
) -> None:
|
||||
self._llm_client = llm_client
|
||||
self._visual_analysis_model = visual_analysis_model
|
||||
self._visual_analysis_timeout = visual_analysis_timeout
|
||||
|
||||
async def analyze_tool_result(
|
||||
self,
|
||||
result: ToolResult,
|
||||
tool_name: str,
|
||||
tool_call: Dict,
|
||||
backend: str,
|
||||
task_description: str = "",
|
||||
) -> ToolResult:
|
||||
"""Callback for LLMClient to handle visual analysis after tool execution."""
|
||||
skip_visual_analysis = False
|
||||
try:
|
||||
arguments = tool_call.function.arguments
|
||||
if isinstance(arguments, str):
|
||||
args = json.loads(arguments.strip() or "{}")
|
||||
else:
|
||||
args = arguments
|
||||
|
||||
if isinstance(args, dict) and args.get("skip_visual_analysis"):
|
||||
skip_visual_analysis = True
|
||||
logger.info(f"Visual analysis skipped for {tool_name} (meta-parameter set by LLM)")
|
||||
except Exception as e:
|
||||
logger.debug(f"Could not parse tool arguments: {e}")
|
||||
|
||||
if skip_visual_analysis:
|
||||
return result
|
||||
|
||||
if backend != "gui":
|
||||
return result
|
||||
|
||||
metadata = getattr(result, "metadata", None)
|
||||
has_screenshots = metadata and (
|
||||
metadata.get("screenshot") or metadata.get("screenshots")
|
||||
)
|
||||
|
||||
if not has_screenshots:
|
||||
try:
|
||||
logger.info(f"No visual data from {tool_name}, capturing screenshot...")
|
||||
screenshot_client = ScreenshotClient()
|
||||
screenshot_bytes = await screenshot_client.capture()
|
||||
|
||||
if screenshot_bytes:
|
||||
if metadata is None:
|
||||
result.metadata = {}
|
||||
metadata = result.metadata
|
||||
metadata["screenshot"] = screenshot_bytes
|
||||
has_screenshots = True
|
||||
logger.info("Screenshot captured for visual analysis")
|
||||
else:
|
||||
logger.warning("Failed to capture screenshot")
|
||||
except Exception as e:
|
||||
logger.warning(f"Error capturing screenshot: {e}")
|
||||
|
||||
if not has_screenshots:
|
||||
logger.debug(f"No visual data available for {tool_name}")
|
||||
return result
|
||||
|
||||
return await self._enhance_result(result, tool_name, task_description)
|
||||
|
||||
async def _enhance_result(
|
||||
self,
|
||||
result: ToolResult,
|
||||
tool_name: str,
|
||||
task_description: str = "",
|
||||
) -> ToolResult:
|
||||
"""Enhance tool result with LLM-based visual analysis."""
|
||||
import asyncio
|
||||
import base64
|
||||
import litellm
|
||||
|
||||
try:
|
||||
metadata = getattr(result, "metadata", None)
|
||||
if not metadata:
|
||||
return result
|
||||
|
||||
screenshots_bytes: List[bytes] = []
|
||||
if metadata.get("screenshots"):
|
||||
screenshots_list = metadata["screenshots"]
|
||||
if isinstance(screenshots_list, list):
|
||||
screenshots_bytes = [s for s in screenshots_list if s]
|
||||
elif metadata.get("screenshot"):
|
||||
screenshots_bytes = [metadata["screenshot"]]
|
||||
|
||||
if not screenshots_bytes:
|
||||
return result
|
||||
|
||||
selected_screenshots = self._select_key_screenshots(
|
||||
screenshots_bytes, max_count=3
|
||||
)
|
||||
|
||||
visual_b64_list = []
|
||||
for visual_data in selected_screenshots:
|
||||
if isinstance(visual_data, bytes):
|
||||
visual_b64_list.append(
|
||||
base64.b64encode(visual_data).decode("utf-8")
|
||||
)
|
||||
else:
|
||||
visual_b64_list.append(visual_data)
|
||||
|
||||
num_screenshots = len(visual_b64_list)
|
||||
|
||||
prompt = GroundingAgentPrompts.visual_analysis(
|
||||
tool_name=tool_name,
|
||||
num_screenshots=num_screenshots,
|
||||
task_description=task_description,
|
||||
)
|
||||
|
||||
content: List[Dict[str, Any]] = [{"type": "text", "text": prompt}]
|
||||
for visual_b64 in visual_b64_list:
|
||||
content.append(
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": f"data:image/png;base64,{visual_b64}"},
|
||||
}
|
||||
)
|
||||
|
||||
visual_model = self._visual_analysis_model or (
|
||||
self._llm_client.model
|
||||
if self._llm_client
|
||||
else "openrouter/anthropic/claude-sonnet-4.5"
|
||||
)
|
||||
_llm_extra: Dict[str, Any] = {}
|
||||
if self._llm_client and visual_model == self._llm_client.model:
|
||||
_llm_extra = (
|
||||
getattr(self._llm_client, "litellm_kwargs", {}) or {}
|
||||
)
|
||||
elif self._visual_analysis_model:
|
||||
try:
|
||||
from openspace.host_detection import build_llm_kwargs
|
||||
|
||||
visual_model, _llm_extra = build_llm_kwargs(visual_model)
|
||||
except Exception as e:
|
||||
logger.debug(
|
||||
f"Failed to resolve dedicated visual model credentials: {e}"
|
||||
)
|
||||
_llm_extra = {}
|
||||
|
||||
response = await asyncio.wait_for(
|
||||
litellm.acompletion(
|
||||
model=visual_model,
|
||||
messages=[{"role": "user", "content": content}],
|
||||
timeout=self._visual_analysis_timeout,
|
||||
**_llm_extra,
|
||||
),
|
||||
timeout=self._visual_analysis_timeout + 5,
|
||||
)
|
||||
|
||||
analysis = response.choices[0].message.content.strip()
|
||||
|
||||
original_content = result.content or "(no text output)"
|
||||
enhanced_content = (
|
||||
f"{original_content}\n\n**Visual content**: {analysis}"
|
||||
)
|
||||
|
||||
enhanced_result = ToolResult(
|
||||
status=result.status,
|
||||
content=enhanced_content,
|
||||
error=result.error,
|
||||
metadata={
|
||||
**metadata,
|
||||
"visual_analyzed": True,
|
||||
"visual_analysis": analysis,
|
||||
},
|
||||
execution_time=result.execution_time,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
f"Enhanced {tool_name} result with visual analysis "
|
||||
f"({num_screenshots} screenshot(s))"
|
||||
)
|
||||
return enhanced_result
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning(
|
||||
f"Visual analysis timed out for {tool_name}, returning original result"
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Failed to analyze visual content for {tool_name}: {e}"
|
||||
)
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _select_key_screenshots(
|
||||
screenshots: List[bytes],
|
||||
max_count: int = 3,
|
||||
) -> List[bytes]:
|
||||
"""Select key screenshots from a sequence, preferring first/last/evenly-spaced."""
|
||||
if len(screenshots) <= max_count:
|
||||
return screenshots
|
||||
|
||||
selected_indices: set[int] = set()
|
||||
|
||||
selected_indices.add(len(screenshots) - 1)
|
||||
|
||||
if max_count >= 2:
|
||||
selected_indices.add(0)
|
||||
|
||||
remaining_slots = max_count - len(selected_indices)
|
||||
if remaining_slots > 0:
|
||||
available_indices = [
|
||||
i
|
||||
for i in range(1, len(screenshots) - 1)
|
||||
if i not in selected_indices
|
||||
]
|
||||
|
||||
if available_indices:
|
||||
step = max(1, len(available_indices) // (remaining_slots + 1))
|
||||
for i in range(remaining_slots):
|
||||
idx = min((i + 1) * step, len(available_indices) - 1)
|
||||
if idx < len(available_indices):
|
||||
selected_indices.add(available_indices[idx])
|
||||
|
||||
selected = [screenshots[i] for i in sorted(selected_indices)]
|
||||
|
||||
logger.debug(
|
||||
f"Selected {len(selected)} screenshots at indices "
|
||||
f"{sorted(selected_indices)} from total of {len(screenshots)}"
|
||||
)
|
||||
|
||||
return selected
|
||||
27
openspace/communication/__init__.py
Normal file
27
openspace/communication/__init__.py
Normal file
|
|
@ -0,0 +1,27 @@
|
|||
from openspace.communication.config import CommunicationConfig, load_communication_config
|
||||
from openspace.communication.session_store import SessionStore, build_session_key
|
||||
from openspace.communication.types import (
|
||||
AttachmentKind,
|
||||
ChannelAttachment,
|
||||
ChannelMessage,
|
||||
ChannelPlatform,
|
||||
ChannelReply,
|
||||
ChannelSession,
|
||||
ChannelSource,
|
||||
SendResult,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AttachmentKind",
|
||||
"ChannelAttachment",
|
||||
"ChannelMessage",
|
||||
"ChannelPlatform",
|
||||
"ChannelReply",
|
||||
"ChannelSession",
|
||||
"ChannelSource",
|
||||
"CommunicationConfig",
|
||||
"SendResult",
|
||||
"SessionStore",
|
||||
"build_session_key",
|
||||
"load_communication_config",
|
||||
]
|
||||
9
openspace/communication/adapters/__init__.py
Normal file
9
openspace/communication/adapters/__init__.py
Normal file
|
|
@ -0,0 +1,9 @@
|
|||
from openspace.communication.adapters.base import BaseChannelAdapter
|
||||
from openspace.communication.adapters.feishu import FeishuAdapter
|
||||
from openspace.communication.adapters.whatsapp import WhatsAppAdapter
|
||||
|
||||
__all__ = [
|
||||
"BaseChannelAdapter",
|
||||
"FeishuAdapter",
|
||||
"WhatsAppAdapter",
|
||||
]
|
||||
63
openspace/communication/adapters/base.py
Normal file
63
openspace/communication/adapters/base.py
Normal file
|
|
@ -0,0 +1,63 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Awaitable, Callable, Optional
|
||||
|
||||
from openspace.utils.logging import Logger
|
||||
|
||||
from openspace.communication.types import ChannelMessage, ChannelPlatform, SendResult
|
||||
|
||||
logger = Logger.get_logger(__name__)
|
||||
|
||||
MessageHandler = Callable[[ChannelMessage], Awaitable[None]]
|
||||
|
||||
|
||||
class BaseChannelAdapter(ABC):
|
||||
platform: ChannelPlatform
|
||||
|
||||
def __init__(self, platform: ChannelPlatform):
|
||||
self.platform = platform
|
||||
self._message_handler: Optional[MessageHandler] = None
|
||||
self._connected = False
|
||||
|
||||
@property
|
||||
def is_connected(self) -> bool:
|
||||
return self._connected
|
||||
|
||||
def set_message_handler(self, handler: MessageHandler) -> None:
|
||||
self._message_handler = handler
|
||||
|
||||
async def dispatch_message(self, message: ChannelMessage) -> None:
|
||||
if self._message_handler is None:
|
||||
logger.warning("Dropping %s message because no handler is attached", self.platform.value)
|
||||
return
|
||||
await self._message_handler(message)
|
||||
|
||||
def register_http_routes(self, app: Any) -> None:
|
||||
"""Optional hook for adapters that need inbound HTTP routes."""
|
||||
|
||||
def validate_configuration(self) -> None:
|
||||
"""Optional hook for adapter-specific startup validation."""
|
||||
|
||||
def get_lock_identity(self) -> Optional[tuple[str, str]]:
|
||||
"""Return an optional (scope, identity) tuple for gateway-scoped locking."""
|
||||
return None
|
||||
|
||||
@abstractmethod
|
||||
async def connect(self) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def disconnect(self) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def send_text(
|
||||
self,
|
||||
chat_id: str,
|
||||
content: str,
|
||||
*,
|
||||
reply_to_message_id: Optional[str] = None,
|
||||
metadata: Optional[dict[str, Any]] = None,
|
||||
) -> SendResult:
|
||||
raise NotImplementedError
|
||||
901
openspace/communication/adapters/feishu.py
Normal file
901
openspace/communication/adapters/feishu.py
Normal file
|
|
@ -0,0 +1,901 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import threading
|
||||
import time
|
||||
from collections import OrderedDict, deque
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
from openspace.communication.adapters.base import BaseChannelAdapter
|
||||
from openspace.communication.attachment_cache import AttachmentCache
|
||||
from openspace.communication.config import FeishuConfig
|
||||
from openspace.communication.policy import is_authorized
|
||||
from openspace.communication.types import (
|
||||
AttachmentKind,
|
||||
ChannelMessage,
|
||||
ChannelPlatform,
|
||||
ChannelSource,
|
||||
SendResult,
|
||||
)
|
||||
from openspace.utils.logging import Logger
|
||||
|
||||
logger = Logger.get_logger(__name__)
|
||||
|
||||
_FEISHU_WEBHOOK_MAX_BODY_BYTES = 1 * 1024 * 1024
|
||||
_FEISHU_WEBHOOK_READ_TIMEOUT_SECONDS = 30
|
||||
_FEISHU_WEBHOOK_RATE_WINDOW_SECONDS = 60
|
||||
_FEISHU_WEBHOOK_RATE_LIMIT_MAX = 120
|
||||
_FEISHU_WEBHOOK_RATE_MAX_KEYS = 4096
|
||||
_FEISHU_WEBHOOK_ANOMALY_TTL_SECONDS = 6 * 60 * 60
|
||||
_FEISHU_DEDUP_CACHE_SIZE = 2048
|
||||
_FEISHU_DEDUP_TTL_SECONDS = 24 * 60 * 60
|
||||
|
||||
try:
|
||||
import lark_oapi as lark
|
||||
from lark_oapi.api.im.v1 import (
|
||||
CreateMessageRequest,
|
||||
CreateMessageRequestBody,
|
||||
GetMessageRequest,
|
||||
GetMessageResourceRequest,
|
||||
ReplyMessageRequest,
|
||||
ReplyMessageRequestBody,
|
||||
)
|
||||
from lark_oapi.core.const import FEISHU_DOMAIN, LARK_DOMAIN
|
||||
|
||||
FEISHU_AVAILABLE = True
|
||||
except ImportError:
|
||||
FEISHU_AVAILABLE = False
|
||||
lark = None # type: ignore[assignment]
|
||||
CreateMessageRequest = None # type: ignore[assignment]
|
||||
CreateMessageRequestBody = None # type: ignore[assignment]
|
||||
GetMessageRequest = None # type: ignore[assignment]
|
||||
GetMessageResourceRequest = None # type: ignore[assignment]
|
||||
ReplyMessageRequest = None # type: ignore[assignment]
|
||||
ReplyMessageRequestBody = None # type: ignore[assignment]
|
||||
FEISHU_DOMAIN = None # type: ignore[assignment]
|
||||
LARK_DOMAIN = None # type: ignore[assignment]
|
||||
|
||||
|
||||
class FeishuAdapter(BaseChannelAdapter):
|
||||
MAX_MESSAGE_LENGTH = 8000
|
||||
_REPLY_CONTEXT_MAX_LEN = 200
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: FeishuConfig,
|
||||
attachment_cache: AttachmentCache,
|
||||
*,
|
||||
runtime_dir: Optional[Path] = None,
|
||||
):
|
||||
super().__init__(ChannelPlatform.FEISHU)
|
||||
self.config = config
|
||||
self.attachment_cache = attachment_cache
|
||||
self.runtime_dir = (
|
||||
Path(runtime_dir).expanduser().resolve()
|
||||
if runtime_dir is not None
|
||||
else attachment_cache.base_dir.parent.resolve()
|
||||
)
|
||||
self._client: Any = None
|
||||
self._bot_open_id = config.bot_open_id
|
||||
self._loop: Optional[asyncio.AbstractEventLoop] = None
|
||||
self._ws_client: Any = None
|
||||
self._ws_thread: Optional[threading.Thread] = None
|
||||
self._running = False
|
||||
self._dedup_state_path = self.runtime_dir / "feishu_seen_message_ids.json"
|
||||
self._seen_message_ids: OrderedDict[str, float] = OrderedDict()
|
||||
self._recent_sent_message_ids: OrderedDict[str, None] = OrderedDict()
|
||||
self._rate_windows: dict[str, deque[float]] = {}
|
||||
self._webhook_anomalies: dict[str, tuple[int, str, float]] = {}
|
||||
self._dedup_dirty = False
|
||||
self._load_seen_message_ids()
|
||||
|
||||
def register_http_routes(self, app: Any) -> None:
|
||||
if self.config.connection_mode == "webhook":
|
||||
app.router.add_post(self.config.webhook_path, self._handle_webhook)
|
||||
|
||||
def validate_configuration(self) -> None:
|
||||
if self.config.connection_mode == "webhook" and not _optional_str(self.config.verification_token):
|
||||
raise ValueError("Feishu webhook mode requires verification_token")
|
||||
|
||||
def get_lock_identity(self) -> Optional[tuple[str, str]]:
|
||||
app_id = _optional_str(self.config.app_id)
|
||||
if not app_id:
|
||||
return None
|
||||
return ("feishu-app", app_id)
|
||||
|
||||
async def connect(self) -> bool:
|
||||
self.validate_configuration()
|
||||
if not FEISHU_AVAILABLE:
|
||||
logger.error("Feishu adapter requires lark-oapi")
|
||||
return False
|
||||
if not self.config.app_id or not self.config.app_secret:
|
||||
logger.error("Feishu adapter missing app_id/app_secret")
|
||||
return False
|
||||
domain = FEISHU_DOMAIN if self.config.domain != "lark" else LARK_DOMAIN
|
||||
self._client = (
|
||||
lark.Client.builder()
|
||||
.app_id(self.config.app_id)
|
||||
.app_secret(self.config.app_secret)
|
||||
.domain(domain)
|
||||
.log_level(lark.LogLevel.WARNING)
|
||||
.build()
|
||||
)
|
||||
self._running = True
|
||||
self._loop = asyncio.get_running_loop()
|
||||
if not self._bot_open_id:
|
||||
self._bot_open_id = await asyncio.to_thread(self._fetch_bot_open_id)
|
||||
if self.config.connection_mode == "websocket":
|
||||
self._start_websocket_client()
|
||||
self._connected = False
|
||||
else:
|
||||
self._connected = True
|
||||
logger.info(
|
||||
"Feishu adapter connected via %s mode",
|
||||
self.config.connection_mode,
|
||||
)
|
||||
return True
|
||||
|
||||
async def disconnect(self) -> None:
|
||||
self._running = False
|
||||
self._connected = False
|
||||
if self._ws_thread is not None and self._ws_thread.is_alive():
|
||||
await asyncio.to_thread(self._ws_thread.join, 5)
|
||||
self._ws_thread = None
|
||||
self._persist_seen_message_ids()
|
||||
self._client = None
|
||||
|
||||
async def send_text(
|
||||
self,
|
||||
chat_id: str,
|
||||
content: str,
|
||||
*,
|
||||
reply_to_message_id: Optional[str] = None,
|
||||
metadata: Optional[dict[str, Any]] = None,
|
||||
) -> SendResult:
|
||||
if not self._client:
|
||||
return SendResult(success=False, error="Feishu client not initialized")
|
||||
|
||||
last_message_id: Optional[str] = None
|
||||
for chunk in _split_text(content, self.MAX_MESSAGE_LENGTH):
|
||||
payload = json.dumps({"text": chunk}, ensure_ascii=False)
|
||||
if reply_to_message_id:
|
||||
body = (
|
||||
ReplyMessageRequestBody.builder()
|
||||
.msg_type("text")
|
||||
.content(payload)
|
||||
.build()
|
||||
)
|
||||
request = (
|
||||
ReplyMessageRequest.builder()
|
||||
.message_id(reply_to_message_id)
|
||||
.request_body(body)
|
||||
.build()
|
||||
)
|
||||
response = await asyncio.to_thread(self._client.im.v1.message.reply, request)
|
||||
else:
|
||||
body = (
|
||||
CreateMessageRequestBody.builder()
|
||||
.receive_id(chat_id)
|
||||
.msg_type("text")
|
||||
.content(payload)
|
||||
.build()
|
||||
)
|
||||
request = (
|
||||
CreateMessageRequest.builder()
|
||||
.receive_id_type("chat_id")
|
||||
.request_body(body)
|
||||
.build()
|
||||
)
|
||||
response = await asyncio.to_thread(self._client.im.v1.message.create, request)
|
||||
|
||||
if not response.success():
|
||||
return SendResult(
|
||||
success=False,
|
||||
error=f"[{response.code}] {response.msg}",
|
||||
raw_response=response,
|
||||
)
|
||||
last_message_id = getattr(getattr(response, "data", None), "message_id", None)
|
||||
if last_message_id:
|
||||
self._remember_sent_message_id(last_message_id)
|
||||
return SendResult(success=True, message_id=last_message_id)
|
||||
|
||||
async def _handle_webhook(self, request: web.Request) -> web.Response:
|
||||
remote_ip = _client_ip_from_request(request)
|
||||
rate_key = f"{self.config.app_id}:{self.config.webhook_path}:{remote_ip}"
|
||||
if not self._check_webhook_rate_limit(rate_key):
|
||||
self._record_webhook_anomaly(remote_ip, "429")
|
||||
return web.Response(status=429, text="Rate limit exceeded")
|
||||
|
||||
content_length = request.content_length or 0
|
||||
if content_length > _FEISHU_WEBHOOK_MAX_BODY_BYTES:
|
||||
self._record_webhook_anomaly(remote_ip, "413")
|
||||
return web.Response(status=413, text="Payload too large")
|
||||
|
||||
try:
|
||||
async with asyncio.timeout(_FEISHU_WEBHOOK_READ_TIMEOUT_SECONDS):
|
||||
body_bytes = await request.read()
|
||||
except TimeoutError:
|
||||
self._record_webhook_anomaly(remote_ip, "408")
|
||||
return web.Response(status=408, text="Request timeout")
|
||||
|
||||
if len(body_bytes) > _FEISHU_WEBHOOK_MAX_BODY_BYTES:
|
||||
self._record_webhook_anomaly(remote_ip, "413")
|
||||
return web.Response(status=413, text="Payload too large")
|
||||
|
||||
try:
|
||||
payload = json.loads(body_bytes.decode("utf-8"))
|
||||
except (UnicodeDecodeError, json.JSONDecodeError):
|
||||
self._record_webhook_anomaly(remote_ip, "400")
|
||||
return web.json_response({"code": 400, "msg": "invalid json"}, status=400)
|
||||
|
||||
incoming_token = str((payload.get("header") or {}).get("token") or payload.get("token") or "")
|
||||
if not incoming_token or not hmac.compare_digest(incoming_token, self.config.verification_token or ""):
|
||||
self._record_webhook_anomaly(remote_ip, "401-token")
|
||||
return web.Response(status=401, text="Invalid verification token")
|
||||
|
||||
if self.config.encrypt_key and not _is_webhook_signature_valid(
|
||||
encrypt_key=self.config.encrypt_key,
|
||||
headers=request.headers,
|
||||
body_bytes=body_bytes,
|
||||
):
|
||||
self._record_webhook_anomaly(remote_ip, "401-signature")
|
||||
return web.Response(status=401, text="Invalid signature")
|
||||
if payload.get("encrypt"):
|
||||
self._record_webhook_anomaly(remote_ip, "400-encrypted")
|
||||
return web.json_response(
|
||||
{"code": 400, "msg": "encrypted webhook payloads are not supported"},
|
||||
status=400,
|
||||
)
|
||||
|
||||
self._clear_webhook_anomaly(remote_ip)
|
||||
|
||||
if payload.get("type") == "url_verification":
|
||||
return web.json_response({"challenge": payload.get("challenge", "")})
|
||||
|
||||
event_type = str((payload.get("header") or {}).get("event_type") or "")
|
||||
if event_type == "im.message.receive_v1":
|
||||
await self._handle_message_event(payload)
|
||||
return web.json_response({"code": 0, "msg": "ok"})
|
||||
|
||||
def _start_websocket_client(self) -> None:
|
||||
assert lark is not None
|
||||
|
||||
handler = (
|
||||
lark.EventDispatcherHandler.builder(
|
||||
self.config.encrypt_key or "",
|
||||
self.config.verification_token or "",
|
||||
)
|
||||
.register_p2_im_message_receive_v1(self._on_message_sync)
|
||||
.build()
|
||||
)
|
||||
self._ws_client = lark.ws.Client(
|
||||
self.config.app_id,
|
||||
self.config.app_secret,
|
||||
event_handler=handler,
|
||||
log_level=lark.LogLevel.INFO,
|
||||
)
|
||||
|
||||
def _run_ws_forever() -> None:
|
||||
import lark_oapi.ws.client as lark_ws_client
|
||||
|
||||
ws_loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(ws_loop)
|
||||
lark_ws_client.loop = ws_loop
|
||||
try:
|
||||
while self._running:
|
||||
try:
|
||||
self._ws_client.start()
|
||||
except Exception as exc:
|
||||
self._connected = False
|
||||
logger.warning("Feishu WebSocket client error: %s", exc)
|
||||
if self._running:
|
||||
time.sleep(5)
|
||||
finally:
|
||||
self._connected = False
|
||||
ws_loop.close()
|
||||
|
||||
self._ws_thread = threading.Thread(
|
||||
target=_run_ws_forever,
|
||||
daemon=True,
|
||||
name="openspace-feishu-ws",
|
||||
)
|
||||
self._ws_thread.start()
|
||||
|
||||
def _on_message_sync(self, data: Any) -> None:
|
||||
if not self._loop or not self._running:
|
||||
return
|
||||
self._connected = True
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
self._handle_websocket_message(data),
|
||||
self._loop,
|
||||
)
|
||||
|
||||
async def _handle_message_event(self, payload: dict[str, Any]) -> None:
|
||||
normalized = await self._normalize_webhook_payload(payload)
|
||||
if normalized is not None:
|
||||
await self.dispatch_message(normalized)
|
||||
|
||||
async def _handle_websocket_message(self, data: Any) -> None:
|
||||
normalized = await self._normalize_websocket_event(data)
|
||||
if normalized is not None:
|
||||
await self.dispatch_message(normalized)
|
||||
|
||||
async def _normalize_webhook_payload(self, payload: dict[str, Any]) -> Optional[ChannelMessage]:
|
||||
event = payload.get("event") or {}
|
||||
message = event.get("message") or {}
|
||||
sender = event.get("sender") or {}
|
||||
sender_id = sender.get("sender_id") or {}
|
||||
if str(sender.get("sender_type", "")).lower() == "bot":
|
||||
return None
|
||||
return await self._normalize_inbound_message(
|
||||
message_id=_optional_str(message.get("message_id")),
|
||||
chat_id=_optional_str(message.get("chat_id")),
|
||||
chat_type=_optional_str(message.get("chat_type")) or "p2p",
|
||||
sender_uid=(
|
||||
_optional_str(sender_id.get("open_id"))
|
||||
or _optional_str(sender_id.get("user_id"))
|
||||
or _optional_str(sender_id.get("union_id"))
|
||||
),
|
||||
sender_name=(
|
||||
_optional_str(sender_id.get("name"))
|
||||
or _optional_str(sender.get("sender_name"))
|
||||
),
|
||||
thread_id=_optional_str(message.get("thread_id")),
|
||||
message_type=_optional_str(message.get("message_type")) or "",
|
||||
content=_safe_json_loads(message.get("content", "{}")),
|
||||
mentions=list(message.get("mentions") or []),
|
||||
reply_to_message_id=(
|
||||
_optional_str(message.get("parent_id"))
|
||||
or _optional_str(message.get("upper_message_id"))
|
||||
),
|
||||
metadata={"webhook_payload": payload},
|
||||
resolve_mentions=False,
|
||||
)
|
||||
|
||||
async def _normalize_websocket_event(self, data: Any) -> Optional[ChannelMessage]:
|
||||
event = getattr(data, "event", None)
|
||||
message = getattr(event, "message", None)
|
||||
sender = getattr(event, "sender", None)
|
||||
if message is None or sender is None:
|
||||
return None
|
||||
if str(getattr(sender, "sender_type", "")).lower() == "bot":
|
||||
return None
|
||||
|
||||
sender_id = getattr(sender, "sender_id", None)
|
||||
return await self._normalize_inbound_message(
|
||||
message_id=_optional_str(getattr(message, "message_id", None)),
|
||||
chat_id=_optional_str(getattr(message, "chat_id", None)),
|
||||
chat_type=_optional_str(getattr(message, "chat_type", None)) or "p2p",
|
||||
sender_uid=(
|
||||
_optional_str(getattr(sender_id, "open_id", None))
|
||||
or _optional_str(getattr(sender_id, "user_id", None))
|
||||
or _optional_str(getattr(sender_id, "union_id", None))
|
||||
),
|
||||
sender_name=_optional_str(getattr(sender, "sender_name", None)),
|
||||
thread_id=_optional_str(getattr(message, "thread_id", None)),
|
||||
message_type=_optional_str(getattr(message, "message_type", None)) or "",
|
||||
content=_safe_json_loads(getattr(message, "content", "{}")),
|
||||
mentions=list(getattr(message, "mentions", None) or []),
|
||||
reply_to_message_id=(
|
||||
_optional_str(getattr(message, "parent_id", None))
|
||||
or _optional_str(getattr(message, "upper_message_id", None))
|
||||
),
|
||||
metadata={"websocket_event": True},
|
||||
resolve_mentions=True,
|
||||
)
|
||||
|
||||
async def _normalize_inbound_message(
|
||||
self,
|
||||
*,
|
||||
message_id: Optional[str],
|
||||
chat_id: Optional[str],
|
||||
chat_type: str,
|
||||
sender_uid: Optional[str],
|
||||
sender_name: Optional[str],
|
||||
thread_id: Optional[str],
|
||||
message_type: str,
|
||||
content: dict[str, Any],
|
||||
mentions: list[Any],
|
||||
reply_to_message_id: Optional[str],
|
||||
metadata: dict[str, Any],
|
||||
resolve_mentions: bool,
|
||||
) -> Optional[ChannelMessage]:
|
||||
if not message_id or not chat_id:
|
||||
return None
|
||||
if self._is_message_seen(message_id):
|
||||
logger.debug("Skipping duplicate Feishu message %s", message_id)
|
||||
return None
|
||||
|
||||
source = ChannelSource(
|
||||
platform=ChannelPlatform.FEISHU,
|
||||
chat_id=chat_id,
|
||||
chat_type="dm" if str(chat_type).lower() == "p2p" else "group",
|
||||
user_id=sender_uid,
|
||||
user_name=sender_name,
|
||||
chat_name=chat_id,
|
||||
thread_id=thread_id,
|
||||
)
|
||||
session_key = _build_session_key_hint(source)
|
||||
normalized_type = str(message_type or "").strip().lower()
|
||||
mentions_bot = self._mentions_bot(mentions)
|
||||
|
||||
text = ""
|
||||
if normalized_type == "text":
|
||||
text = str(content.get("text", "")).strip()
|
||||
if resolve_mentions:
|
||||
text = _resolve_mentions(text, mentions)
|
||||
elif normalized_type == "post":
|
||||
text = _extract_post_text(content)
|
||||
|
||||
prefilter_message = ChannelMessage(
|
||||
source=source,
|
||||
text=text,
|
||||
message_id=message_id,
|
||||
reply_to_message_id=reply_to_message_id,
|
||||
mentions_bot=mentions_bot,
|
||||
metadata=metadata,
|
||||
)
|
||||
if not self._passes_prefilter(prefilter_message):
|
||||
self._remember_message_seen(message_id)
|
||||
return None
|
||||
|
||||
attachments = []
|
||||
if normalized_type == "image":
|
||||
attachment = await self._download_attachment(
|
||||
session_key=session_key,
|
||||
message_id=message_id,
|
||||
file_key=str(content.get("image_key", "")).strip(),
|
||||
file_name=str(content.get("image_key", "image")).strip() + ".png",
|
||||
kind=AttachmentKind.IMAGE,
|
||||
resource_type="image",
|
||||
)
|
||||
if attachment is not None:
|
||||
attachments.append(attachment)
|
||||
elif normalized_type == "file":
|
||||
attachment = await self._download_attachment(
|
||||
session_key=session_key,
|
||||
message_id=message_id,
|
||||
file_key=str(content.get("file_key", "")).strip(),
|
||||
file_name=str(content.get("file_name", "document")).strip(),
|
||||
kind=AttachmentKind.DOCUMENT,
|
||||
resource_type="file",
|
||||
)
|
||||
if attachment is not None:
|
||||
attachments.append(attachment)
|
||||
|
||||
prefilter_message.attachments = attachments
|
||||
prefilter_message.reply_to_text = await self._fetch_message_text(reply_to_message_id)
|
||||
self._remember_message_seen(message_id)
|
||||
return prefilter_message
|
||||
|
||||
def _passes_prefilter(self, message: ChannelMessage) -> bool:
|
||||
if not is_authorized(message, self.config):
|
||||
logger.info("Rejected Feishu message from unauthorized user %s", message.source.user_id)
|
||||
return False
|
||||
if message.source.chat_type == "dm":
|
||||
return self.config.allow_dm
|
||||
if not self.config.allow_groups:
|
||||
return False
|
||||
if self.config.group_policy == "disabled":
|
||||
return False
|
||||
if self.config.group_policy == "mention_only":
|
||||
return message.mentions_bot
|
||||
if self.config.group_policy == "reply_or_mention":
|
||||
return message.mentions_bot or self._is_reply_to_recent_bot_message(
|
||||
message.reply_to_message_id
|
||||
)
|
||||
return True
|
||||
|
||||
def _is_reply_to_recent_bot_message(self, message_id: Optional[str]) -> bool:
|
||||
if not message_id:
|
||||
return False
|
||||
return message_id in self._recent_sent_message_ids
|
||||
|
||||
def _remember_sent_message_id(self, message_id: str) -> None:
|
||||
self._recent_sent_message_ids.pop(message_id, None)
|
||||
self._recent_sent_message_ids[message_id] = None
|
||||
while len(self._recent_sent_message_ids) > _FEISHU_DEDUP_CACHE_SIZE:
|
||||
self._recent_sent_message_ids.popitem(last=False)
|
||||
|
||||
async def _download_attachment(
|
||||
self,
|
||||
*,
|
||||
session_key: str,
|
||||
message_id: str,
|
||||
file_key: str,
|
||||
file_name: str,
|
||||
kind: AttachmentKind,
|
||||
resource_type: str,
|
||||
):
|
||||
if not self._client or not file_key:
|
||||
return None
|
||||
request = (
|
||||
GetMessageResourceRequest.builder()
|
||||
.message_id(message_id)
|
||||
.file_key(file_key)
|
||||
.type(resource_type)
|
||||
.build()
|
||||
)
|
||||
response = await asyncio.to_thread(self._client.im.v1.message_resource.get, request)
|
||||
if not response.success():
|
||||
logger.warning(
|
||||
"Failed to download Feishu attachment: code=%s msg=%s",
|
||||
response.code,
|
||||
response.msg,
|
||||
)
|
||||
return None
|
||||
try:
|
||||
file_data = await asyncio.to_thread(
|
||||
_read_attachment_body,
|
||||
response.file,
|
||||
self.attachment_cache.max_attachment_bytes,
|
||||
)
|
||||
except ValueError as exc:
|
||||
logger.warning(
|
||||
"Rejected Feishu attachment for session %s: %s",
|
||||
session_key,
|
||||
exc,
|
||||
)
|
||||
return None
|
||||
return self.attachment_cache.save_bytes(
|
||||
session_key=session_key,
|
||||
data=file_data,
|
||||
filename=file_name,
|
||||
kind=kind,
|
||||
)
|
||||
|
||||
def _fetch_bot_open_id(self) -> Optional[str]:
|
||||
if not self._client or not lark:
|
||||
return None
|
||||
try:
|
||||
request = (
|
||||
lark.BaseRequest.builder()
|
||||
.http_method(lark.HttpMethod.GET)
|
||||
.uri("/open-apis/bot/v3/info")
|
||||
.token_types({lark.AccessTokenType.APP})
|
||||
.build()
|
||||
)
|
||||
response = self._client.request(request)
|
||||
if not response.success():
|
||||
logger.warning(
|
||||
"Failed to fetch Feishu bot info: code=%s msg=%s",
|
||||
response.code,
|
||||
response.msg,
|
||||
)
|
||||
return None
|
||||
payload = json.loads(response.raw.content)
|
||||
bot = (payload.get("data") or payload).get("bot") or {}
|
||||
return _optional_str(bot.get("open_id"))
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to resolve Feishu bot open_id: %s", exc)
|
||||
return None
|
||||
|
||||
async def _fetch_message_text(self, message_id: Optional[str]) -> Optional[str]:
|
||||
if not self._client or not message_id or GetMessageRequest is None:
|
||||
return None
|
||||
|
||||
request = GetMessageRequest.builder().message_id(message_id).build()
|
||||
try:
|
||||
response = await asyncio.to_thread(self._client.im.v1.message.get, request)
|
||||
except Exception as exc:
|
||||
logger.debug("Failed to fetch Feishu parent message %s: %s", message_id, exc)
|
||||
return None
|
||||
if not response.success():
|
||||
return None
|
||||
|
||||
data = getattr(response, "data", None)
|
||||
message_obj = None
|
||||
items = getattr(data, "items", None)
|
||||
if items:
|
||||
message_obj = items[0]
|
||||
elif data is not None:
|
||||
message_obj = getattr(data, "message", None) or data
|
||||
if message_obj is None:
|
||||
return None
|
||||
|
||||
body = getattr(message_obj, "body", None)
|
||||
raw_content = getattr(body, "content", None) if body is not None else getattr(message_obj, "content", None)
|
||||
message_type = (
|
||||
getattr(message_obj, "msg_type", None)
|
||||
or getattr(message_obj, "message_type", None)
|
||||
or ""
|
||||
)
|
||||
content = _safe_json_loads(raw_content)
|
||||
text = ""
|
||||
if str(message_type).lower() == "text":
|
||||
text = str(content.get("text", "")).strip()
|
||||
elif str(message_type).lower() == "post":
|
||||
text = _extract_post_text(content)
|
||||
if not text:
|
||||
return None
|
||||
if len(text) > self._REPLY_CONTEXT_MAX_LEN:
|
||||
text = text[: self._REPLY_CONTEXT_MAX_LEN] + "..."
|
||||
return text
|
||||
|
||||
def _mentions_bot(self, mentions: list[Any]) -> bool:
|
||||
if not mentions:
|
||||
return False
|
||||
if not self._bot_open_id:
|
||||
return False
|
||||
|
||||
for mention in mentions:
|
||||
mention_id = (mention.get("id") or {}) if isinstance(mention, dict) else getattr(mention, "id", None)
|
||||
open_id = (
|
||||
_optional_str(mention_id.get("open_id"))
|
||||
if isinstance(mention_id, dict)
|
||||
else _optional_str(getattr(mention_id, "open_id", None))
|
||||
)
|
||||
if open_id == self._bot_open_id:
|
||||
return True
|
||||
return False
|
||||
|
||||
def _load_seen_message_ids(self) -> None:
|
||||
if not self._dedup_state_path.exists():
|
||||
return
|
||||
try:
|
||||
payload = json.loads(self._dedup_state_path.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError):
|
||||
logger.warning("Failed to load Feishu dedup cache from %s", self._dedup_state_path)
|
||||
return
|
||||
entries = payload.get("message_ids", {}) if isinstance(payload, dict) else {}
|
||||
now = time.time()
|
||||
valid: list[tuple[str, float]] = []
|
||||
if isinstance(entries, dict):
|
||||
for message_id, seen_at in entries.items():
|
||||
normalized_id = _optional_str(message_id)
|
||||
if not normalized_id:
|
||||
continue
|
||||
try:
|
||||
timestamp = float(seen_at)
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
if now - timestamp <= _FEISHU_DEDUP_TTL_SECONDS:
|
||||
valid.append((normalized_id, timestamp))
|
||||
for message_id, seen_at in sorted(valid, key=lambda item: item[1])[-_FEISHU_DEDUP_CACHE_SIZE:]:
|
||||
self._seen_message_ids[message_id] = seen_at
|
||||
|
||||
def _persist_seen_message_ids(self) -> None:
|
||||
if not self._dedup_dirty:
|
||||
return
|
||||
self._dedup_state_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
payload = {"message_ids": dict(self._seen_message_ids)}
|
||||
try:
|
||||
self._dedup_state_path.write_text(
|
||||
json.dumps(payload, ensure_ascii=False, indent=2),
|
||||
encoding="utf-8",
|
||||
)
|
||||
except OSError:
|
||||
logger.warning("Failed to persist Feishu dedup cache to %s", self._dedup_state_path)
|
||||
return
|
||||
self._dedup_dirty = False
|
||||
|
||||
def _is_message_seen(self, message_id: str) -> bool:
|
||||
now = time.time()
|
||||
self._prune_seen_message_ids(now)
|
||||
return message_id in self._seen_message_ids
|
||||
|
||||
def _remember_message_seen(self, message_id: str) -> None:
|
||||
now = time.time()
|
||||
self._prune_seen_message_ids(now)
|
||||
if message_id in self._seen_message_ids:
|
||||
self._seen_message_ids.move_to_end(message_id)
|
||||
return
|
||||
self._seen_message_ids[message_id] = now
|
||||
self._seen_message_ids.move_to_end(message_id)
|
||||
while len(self._seen_message_ids) > _FEISHU_DEDUP_CACHE_SIZE:
|
||||
self._seen_message_ids.popitem(last=False)
|
||||
self._dedup_dirty = True
|
||||
self._persist_seen_message_ids()
|
||||
|
||||
def _mark_message_seen(self, message_id: str) -> bool:
|
||||
if self._is_message_seen(message_id):
|
||||
return True
|
||||
self._remember_message_seen(message_id)
|
||||
return False
|
||||
|
||||
def _prune_seen_message_ids(self, now: Optional[float] = None) -> None:
|
||||
current = now or time.time()
|
||||
stale = [
|
||||
message_id
|
||||
for message_id, seen_at in self._seen_message_ids.items()
|
||||
if current - seen_at > _FEISHU_DEDUP_TTL_SECONDS
|
||||
]
|
||||
for message_id in stale:
|
||||
self._seen_message_ids.pop(message_id, None)
|
||||
self._dedup_dirty = True
|
||||
|
||||
def _check_webhook_rate_limit(self, rate_key: str) -> bool:
|
||||
now = time.time()
|
||||
window = self._rate_windows.get(rate_key)
|
||||
if window is None:
|
||||
if len(self._rate_windows) >= _FEISHU_WEBHOOK_RATE_MAX_KEYS:
|
||||
stale_keys = [
|
||||
key
|
||||
for key, timestamps in self._rate_windows.items()
|
||||
if not timestamps or now - timestamps[-1] > _FEISHU_WEBHOOK_RATE_WINDOW_SECONDS
|
||||
]
|
||||
for key in stale_keys:
|
||||
self._rate_windows.pop(key, None)
|
||||
if rate_key not in self._rate_windows and len(self._rate_windows) >= _FEISHU_WEBHOOK_RATE_MAX_KEYS:
|
||||
return False
|
||||
window = deque()
|
||||
self._rate_windows[rate_key] = window
|
||||
cutoff = now - _FEISHU_WEBHOOK_RATE_WINDOW_SECONDS
|
||||
while window and window[0] < cutoff:
|
||||
window.popleft()
|
||||
if len(window) >= _FEISHU_WEBHOOK_RATE_LIMIT_MAX:
|
||||
return False
|
||||
window.append(now)
|
||||
return True
|
||||
|
||||
def _record_webhook_anomaly(self, remote_ip: str, status: str) -> None:
|
||||
now = time.time()
|
||||
current = self._webhook_anomalies.get(remote_ip)
|
||||
if current and now - current[2] < _FEISHU_WEBHOOK_ANOMALY_TTL_SECONDS:
|
||||
self._webhook_anomalies[remote_ip] = (current[0] + 1, status, current[2])
|
||||
return
|
||||
self._webhook_anomalies[remote_ip] = (1, status, now)
|
||||
|
||||
def _clear_webhook_anomaly(self, remote_ip: str) -> None:
|
||||
self._webhook_anomalies.pop(remote_ip, None)
|
||||
|
||||
|
||||
def _safe_json_loads(value: Any) -> dict[str, Any]:
|
||||
if isinstance(value, dict):
|
||||
return value
|
||||
try:
|
||||
return json.loads(str(value or "{}"))
|
||||
except json.JSONDecodeError:
|
||||
return {}
|
||||
|
||||
|
||||
def _optional_str(value: Any) -> Optional[str]:
|
||||
if value is None:
|
||||
return None
|
||||
value = str(value).strip()
|
||||
return value or None
|
||||
|
||||
|
||||
def _client_ip_from_request(request: web.Request) -> str:
|
||||
forwarded = str(request.headers.get("x-forwarded-for", "") or "").split(",")[0].strip()
|
||||
if forwarded:
|
||||
return forwarded
|
||||
peer = request.transport.get_extra_info("peername") if request.transport else None
|
||||
if isinstance(peer, tuple) and peer:
|
||||
return str(peer[0])
|
||||
return request.remote or "unknown"
|
||||
|
||||
|
||||
def _resolve_mentions(text: str, mentions: list[Any]) -> str:
|
||||
if not text or not mentions:
|
||||
return text
|
||||
|
||||
resolved = text
|
||||
for mention in mentions:
|
||||
key = _optional_str(
|
||||
mention.get("key") if isinstance(mention, dict) else getattr(mention, "key", None)
|
||||
)
|
||||
if not key or key not in resolved:
|
||||
continue
|
||||
name = _optional_str(
|
||||
mention.get("name") if isinstance(mention, dict) else getattr(mention, "name", None)
|
||||
) or "user"
|
||||
resolved = resolved.replace(key, f"@{name}")
|
||||
return resolved
|
||||
|
||||
|
||||
def _split_text(content: str, limit: int) -> list[str]:
|
||||
text = content.strip()
|
||||
if not text:
|
||||
return [""]
|
||||
if len(text) <= limit:
|
||||
return [text]
|
||||
chunks = []
|
||||
remaining = text
|
||||
while remaining:
|
||||
chunk = remaining[:limit]
|
||||
if len(remaining) > limit:
|
||||
split_at = chunk.rfind("\n")
|
||||
if split_at < limit // 3:
|
||||
split_at = chunk.rfind(" ")
|
||||
if split_at >= limit // 3:
|
||||
chunk = chunk[:split_at]
|
||||
chunks.append(chunk.strip())
|
||||
remaining = remaining[len(chunk):].lstrip()
|
||||
return [chunk for chunk in chunks if chunk]
|
||||
|
||||
|
||||
def _is_webhook_signature_valid(*, encrypt_key: str, headers: Any, body_bytes: bytes) -> bool:
|
||||
timestamp = str(headers.get("x-lark-request-timestamp", "") or "")
|
||||
nonce = str(headers.get("x-lark-request-nonce", "") or "")
|
||||
signature = str(headers.get("x-lark-signature", "") or "")
|
||||
if not timestamp or not nonce or not signature:
|
||||
return False
|
||||
content = f"{timestamp}{nonce}{encrypt_key}{body_bytes.decode('utf-8', errors='replace')}"
|
||||
expected = hashlib.sha256(content.encode("utf-8")).hexdigest()
|
||||
return hmac.compare_digest(signature, expected)
|
||||
|
||||
|
||||
def _build_session_key_hint(source: ChannelSource) -> str:
|
||||
parts = [source.platform.value, source.chat_id]
|
||||
if source.thread_id:
|
||||
parts.append(source.thread_id)
|
||||
return "__".join(part.replace("/", "_") for part in parts if part)
|
||||
|
||||
|
||||
def _extract_post_text(content: dict[str, Any]) -> str:
|
||||
texts: list[str] = []
|
||||
|
||||
def _walk(value: Any) -> None:
|
||||
if isinstance(value, dict):
|
||||
title = _optional_str(value.get("title"))
|
||||
if title:
|
||||
texts.append(title)
|
||||
tag = _optional_str(value.get("tag"))
|
||||
if tag == "at":
|
||||
user_name = _optional_str(value.get("user_name")) or "user"
|
||||
texts.append(f"@{user_name}")
|
||||
else:
|
||||
text = _optional_str(value.get("text"))
|
||||
if text:
|
||||
texts.append(text)
|
||||
for nested in value.values():
|
||||
_walk(nested)
|
||||
elif isinstance(value, list):
|
||||
for item in value:
|
||||
_walk(item)
|
||||
|
||||
_walk(content)
|
||||
deduped: list[str] = []
|
||||
for text in texts:
|
||||
if text not in deduped:
|
||||
deduped.append(text)
|
||||
return "\n".join(deduped).strip()
|
||||
|
||||
|
||||
def _read_attachment_body(raw_file: Any, max_bytes: int) -> bytes:
|
||||
if raw_file is None:
|
||||
return b""
|
||||
|
||||
if isinstance(raw_file, bytes):
|
||||
data = raw_file
|
||||
elif isinstance(raw_file, bytearray):
|
||||
data = bytes(raw_file)
|
||||
elif hasattr(raw_file, "read"):
|
||||
chunks: list[bytes] = []
|
||||
total = 0
|
||||
try:
|
||||
while True:
|
||||
chunk = raw_file.read(65536)
|
||||
if not chunk:
|
||||
break
|
||||
if isinstance(chunk, str):
|
||||
chunk = chunk.encode("utf-8")
|
||||
elif isinstance(chunk, bytearray):
|
||||
chunk = bytes(chunk)
|
||||
elif not isinstance(chunk, bytes):
|
||||
chunk = bytes(chunk)
|
||||
total += len(chunk)
|
||||
if total > max_bytes:
|
||||
raise ValueError(
|
||||
f"attachment size {total} exceeds limit {max_bytes}"
|
||||
)
|
||||
chunks.append(chunk)
|
||||
finally:
|
||||
close = getattr(raw_file, "close", None)
|
||||
if callable(close):
|
||||
close()
|
||||
data = b"".join(chunks)
|
||||
else:
|
||||
data = bytes(raw_file)
|
||||
|
||||
if len(data) > max_bytes:
|
||||
raise ValueError(f"attachment size {len(data)} exceeds limit {max_bytes}")
|
||||
return data
|
||||
462
openspace/communication/adapters/whatsapp.py
Normal file
462
openspace/communication/adapters/whatsapp.py
Normal file
|
|
@ -0,0 +1,462 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
import shutil
|
||||
import subprocess
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
import aiohttp
|
||||
|
||||
from openspace.communication.adapters.base import BaseChannelAdapter
|
||||
from openspace.communication.attachment_cache import AttachmentCache
|
||||
from openspace.communication.config import WhatsAppConfig
|
||||
from openspace.communication.types import (
|
||||
AttachmentKind,
|
||||
ChannelMessage,
|
||||
ChannelPlatform,
|
||||
ChannelSource,
|
||||
SendResult,
|
||||
)
|
||||
from openspace.utils.logging import Logger
|
||||
|
||||
logger = Logger.get_logger(__name__)
|
||||
|
||||
|
||||
class WhatsAppAdapter(BaseChannelAdapter):
|
||||
def __init__(
|
||||
self,
|
||||
config: WhatsAppConfig,
|
||||
attachment_cache: AttachmentCache,
|
||||
*,
|
||||
runtime_dir: Optional[Path] = None,
|
||||
poll_interval_seconds: float = 1.0,
|
||||
):
|
||||
super().__init__(ChannelPlatform.WHATSAPP)
|
||||
self.config = config
|
||||
self.attachment_cache = attachment_cache
|
||||
self.runtime_dir = (
|
||||
Path(runtime_dir).expanduser().resolve()
|
||||
if runtime_dir is not None
|
||||
else attachment_cache.base_dir.parent.resolve()
|
||||
)
|
||||
self._poll_interval_seconds = poll_interval_seconds
|
||||
self._http_session: Optional[aiohttp.ClientSession] = None
|
||||
self._ws: Optional[aiohttp.ClientWebSocketResponse] = None
|
||||
self._receiver_task: Optional[asyncio.Task] = None
|
||||
self._bridge_process: Optional[subprocess.Popen] = None
|
||||
self._pending_requests: dict[str, asyncio.Future[dict[str, Any]]] = {}
|
||||
self._auth_event = asyncio.Event()
|
||||
self._status_event = asyncio.Event()
|
||||
self._bridge_state = "disconnected"
|
||||
|
||||
def validate_configuration(self) -> None:
|
||||
if self.config.bridge.enforce_loopback and self.config.bridge.host not in {"127.0.0.1", "localhost"}:
|
||||
raise ValueError("WhatsApp bridge host must stay on loopback")
|
||||
|
||||
def get_lock_identity(self) -> Optional[tuple[str, str]]:
|
||||
return ("whatsapp-session", str(self._session_dir().resolve()))
|
||||
|
||||
async def connect(self) -> bool:
|
||||
self.validate_configuration()
|
||||
if self._http_session is None:
|
||||
self._http_session = aiohttp.ClientSession(
|
||||
timeout=aiohttp.ClientTimeout(total=20),
|
||||
)
|
||||
|
||||
for attempt in range(2):
|
||||
try:
|
||||
await self._open_control_socket()
|
||||
except Exception as exc:
|
||||
logger.info("WhatsApp bridge connection attempt %s failed: %s", attempt + 1, exc)
|
||||
if attempt == 0:
|
||||
await self._start_bridge_process()
|
||||
await asyncio.sleep(1)
|
||||
continue
|
||||
return False
|
||||
break
|
||||
|
||||
for _ in range(20):
|
||||
if self._bridge_state == "connected":
|
||||
self._connected = True
|
||||
return True
|
||||
await asyncio.sleep(1)
|
||||
logger.error("WhatsApp bridge control channel opened but WhatsApp session did not connect")
|
||||
return False
|
||||
|
||||
async def disconnect(self) -> None:
|
||||
self._connected = False
|
||||
self._bridge_state = "disconnected"
|
||||
self._status_event.clear()
|
||||
self._auth_event.clear()
|
||||
|
||||
if self._receiver_task is not None:
|
||||
self._receiver_task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError):
|
||||
await self._receiver_task
|
||||
self._receiver_task = None
|
||||
|
||||
if self._ws is not None:
|
||||
await self._ws.close()
|
||||
self._ws = None
|
||||
|
||||
for future in self._pending_requests.values():
|
||||
if not future.done():
|
||||
future.set_exception(RuntimeError("WhatsApp bridge disconnected"))
|
||||
self._pending_requests.clear()
|
||||
|
||||
if self._http_session is not None:
|
||||
await self._http_session.close()
|
||||
self._http_session = None
|
||||
|
||||
if self._bridge_process is not None and self._bridge_process.poll() is None:
|
||||
self._bridge_process.terminate()
|
||||
try:
|
||||
self._bridge_process.wait(timeout=10)
|
||||
except subprocess.TimeoutExpired:
|
||||
self._bridge_process.kill()
|
||||
self._bridge_process = None
|
||||
|
||||
async def send_text(
|
||||
self,
|
||||
chat_id: str,
|
||||
content: str,
|
||||
*,
|
||||
reply_to_message_id: Optional[str] = None,
|
||||
metadata: Optional[dict[str, Any]] = None,
|
||||
) -> SendResult:
|
||||
if self._ws is None:
|
||||
return SendResult(success=False, error="WhatsApp bridge not initialized")
|
||||
|
||||
last_message_id: Optional[str] = None
|
||||
for chunk in _split_text(content, 60000):
|
||||
try:
|
||||
payload = await self._send_command(
|
||||
{
|
||||
"type": "send",
|
||||
"to": chat_id,
|
||||
"text": chunk,
|
||||
"replyToMessageId": reply_to_message_id,
|
||||
}
|
||||
)
|
||||
except Exception as exc:
|
||||
return SendResult(success=False, error=str(exc))
|
||||
last_message_id = _optional_str(payload.get("messageId")) or last_message_id
|
||||
return SendResult(success=True, message_id=last_message_id)
|
||||
|
||||
async def send_media(
|
||||
self,
|
||||
chat_id: str,
|
||||
*,
|
||||
file_path: str,
|
||||
mimetype: str,
|
||||
caption: Optional[str] = None,
|
||||
file_name: Optional[str] = None,
|
||||
reply_to_message_id: Optional[str] = None,
|
||||
) -> SendResult:
|
||||
if self._ws is None:
|
||||
return SendResult(success=False, error="WhatsApp bridge not initialized")
|
||||
try:
|
||||
payload = await self._send_command(
|
||||
{
|
||||
"type": "send_media",
|
||||
"to": chat_id,
|
||||
"filePath": file_path,
|
||||
"mimetype": mimetype,
|
||||
"caption": caption,
|
||||
"fileName": file_name,
|
||||
"replyToMessageId": reply_to_message_id,
|
||||
}
|
||||
)
|
||||
except Exception as exc:
|
||||
return SendResult(success=False, error=str(exc))
|
||||
return SendResult(success=True, message_id=_optional_str(payload.get("messageId")))
|
||||
|
||||
async def _open_control_socket(self) -> None:
|
||||
if self._http_session is None:
|
||||
raise RuntimeError("WhatsApp bridge HTTP session is not initialized")
|
||||
if self._receiver_task is not None:
|
||||
self._receiver_task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError):
|
||||
await self._receiver_task
|
||||
self._receiver_task = None
|
||||
if self._ws is not None:
|
||||
await self._ws.close()
|
||||
self._ws = None
|
||||
self._auth_event.clear()
|
||||
self._status_event.clear()
|
||||
ws = await self._http_session.ws_connect(
|
||||
self.config.bridge.ws_url,
|
||||
heartbeat=20,
|
||||
autoping=True,
|
||||
max_msg_size=4 * 1024 * 1024,
|
||||
)
|
||||
self._ws = ws
|
||||
self._receiver_task = asyncio.create_task(self._receive_loop(ws))
|
||||
await self._send_ws_json({"type": "auth", "token": self._effective_bridge_token()})
|
||||
await asyncio.wait_for(self._auth_event.wait(), timeout=5)
|
||||
|
||||
async def _receive_loop(self, ws: aiohttp.ClientWebSocketResponse) -> None:
|
||||
try:
|
||||
async for msg in ws:
|
||||
if msg.type != aiohttp.WSMsgType.TEXT:
|
||||
if msg.type in {
|
||||
aiohttp.WSMsgType.CLOSE,
|
||||
aiohttp.WSMsgType.CLOSED,
|
||||
aiohttp.WSMsgType.ERROR,
|
||||
}:
|
||||
break
|
||||
continue
|
||||
try:
|
||||
payload = json.loads(msg.data)
|
||||
except json.JSONDecodeError:
|
||||
logger.warning("Ignoring invalid WhatsApp bridge JSON: %r", msg.data[:200])
|
||||
continue
|
||||
await self._handle_ws_payload(payload)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.warning("WhatsApp bridge receive loop stopped: %s", exc)
|
||||
finally:
|
||||
if self._ws is ws:
|
||||
self._ws = None
|
||||
self._connected = False
|
||||
self._bridge_state = "disconnected"
|
||||
self._status_event.clear()
|
||||
for request_id, future in list(self._pending_requests.items()):
|
||||
if not future.done():
|
||||
future.set_exception(RuntimeError("WhatsApp bridge disconnected"))
|
||||
self._pending_requests.pop(request_id, None)
|
||||
|
||||
async def _handle_ws_payload(self, payload: dict[str, Any]) -> None:
|
||||
message_type = str(payload.get("type", "")).strip().lower()
|
||||
if message_type == "auth_ok":
|
||||
self._auth_event.set()
|
||||
return
|
||||
if message_type == "status":
|
||||
self._bridge_state = str(payload.get("status", "")).strip().lower() or "disconnected"
|
||||
self._connected = self._bridge_state == "connected"
|
||||
self._status_event.set()
|
||||
return
|
||||
if message_type == "qr":
|
||||
logger.info("WhatsApp bridge is waiting for QR scan")
|
||||
return
|
||||
if message_type == "ack":
|
||||
request_id = _optional_str(payload.get("requestId"))
|
||||
if request_id and request_id in self._pending_requests:
|
||||
future = self._pending_requests.pop(request_id)
|
||||
if not future.done():
|
||||
future.set_result(payload)
|
||||
return
|
||||
if message_type == "error":
|
||||
request_id = _optional_str(payload.get("requestId"))
|
||||
error = _optional_str(payload.get("error")) or "Unknown bridge error"
|
||||
if request_id and request_id in self._pending_requests:
|
||||
future = self._pending_requests.pop(request_id)
|
||||
if not future.done():
|
||||
future.set_exception(RuntimeError(error))
|
||||
else:
|
||||
logger.warning("WhatsApp bridge error: %s", error)
|
||||
return
|
||||
if message_type == "message":
|
||||
message = await self._normalize_event(payload)
|
||||
if message is not None:
|
||||
await self.dispatch_message(message)
|
||||
|
||||
async def _send_command(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
if self._ws is None:
|
||||
raise RuntimeError("WhatsApp bridge control channel is not connected")
|
||||
request_id = uuid.uuid4().hex[:12]
|
||||
future: asyncio.Future[dict[str, Any]] = asyncio.get_running_loop().create_future()
|
||||
self._pending_requests[request_id] = future
|
||||
try:
|
||||
await self._send_ws_json({**payload, "requestId": request_id})
|
||||
return await asyncio.wait_for(future, timeout=20)
|
||||
finally:
|
||||
self._pending_requests.pop(request_id, None)
|
||||
|
||||
async def _send_ws_json(self, payload: dict[str, Any]) -> None:
|
||||
if self._ws is None:
|
||||
raise RuntimeError("WhatsApp bridge control channel is not connected")
|
||||
await self._ws.send_str(json.dumps(payload, ensure_ascii=False))
|
||||
|
||||
async def _normalize_event(self, event: dict[str, Any]) -> Optional[ChannelMessage]:
|
||||
chat_id = str(event.get("chatId", "")).strip()
|
||||
message_id = str(event.get("messageId", "")).strip()
|
||||
sender_id = str(event.get("senderId", "")).strip()
|
||||
if not chat_id or not message_id:
|
||||
return None
|
||||
normalized_sender_id = _normalize_whatsapp_identifier(sender_id)
|
||||
|
||||
source = ChannelSource(
|
||||
platform=ChannelPlatform.WHATSAPP,
|
||||
chat_id=chat_id,
|
||||
chat_type="group" if event.get("isGroup") else "dm",
|
||||
user_id=normalized_sender_id or sender_id or None,
|
||||
user_name=_optional_str(event.get("senderName")),
|
||||
chat_name=_optional_str(event.get("chatName")),
|
||||
)
|
||||
|
||||
session_key = _build_session_key_hint(source)
|
||||
attachments = []
|
||||
media_type = str(event.get("mediaType", "")).strip().lower()
|
||||
attachment_kind = AttachmentKind.IMAGE if media_type == "image" else AttachmentKind.DOCUMENT
|
||||
for media_path in event.get("mediaUrls") or []:
|
||||
attachment = self.attachment_cache.copy_local_file(
|
||||
session_key=session_key,
|
||||
source_path=str(media_path),
|
||||
kind=attachment_kind,
|
||||
)
|
||||
if attachment is not None:
|
||||
attachments.append(attachment)
|
||||
|
||||
body = str(event.get("body", "") or "").strip()
|
||||
return ChannelMessage(
|
||||
source=source,
|
||||
text=body,
|
||||
message_id=message_id,
|
||||
attachments=attachments,
|
||||
reply_to_message_id=_optional_str(event.get("replyToMessageId")),
|
||||
mentions_bot=bool(event.get("mentionsBot")),
|
||||
metadata={
|
||||
"bridge_event": event,
|
||||
"raw_user_id": sender_id or None,
|
||||
"auth_candidates": [
|
||||
candidate
|
||||
for candidate in (
|
||||
sender_id or None,
|
||||
normalized_sender_id or None,
|
||||
f"+{normalized_sender_id}" if normalized_sender_id else None,
|
||||
)
|
||||
if candidate
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
async def _start_bridge_process(self) -> None:
|
||||
if self._bridge_process is not None and self._bridge_process.poll() is None:
|
||||
return
|
||||
|
||||
bridge_script = self._resolve_bridge_script()
|
||||
bridge_dir = bridge_script.parent
|
||||
session_dir = self._session_dir()
|
||||
session_dir.mkdir(parents=True, exist_ok=True)
|
||||
self._outbound_media_root().mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if self.config.bridge.auto_install_dependencies and not (bridge_dir / "node_modules").exists():
|
||||
subprocess.run(
|
||||
["npm", "install", "--silent"],
|
||||
cwd=bridge_dir,
|
||||
check=True,
|
||||
)
|
||||
|
||||
env = os.environ.copy()
|
||||
env["BRIDGE_TOKEN"] = self._effective_bridge_token()
|
||||
env["BRIDGE_MEDIA_ROOT"] = str(self._outbound_media_root())
|
||||
if self.config.allowed_users:
|
||||
env["WHATSAPP_ALLOWED_USERS"] = ",".join(self.config.allowed_users)
|
||||
if self.config.reply_prefix is not None:
|
||||
env["WHATSAPP_REPLY_PREFIX"] = self.config.reply_prefix
|
||||
|
||||
self._bridge_process = subprocess.Popen(
|
||||
[
|
||||
"node",
|
||||
str(bridge_script),
|
||||
"--host",
|
||||
self.config.bridge.host,
|
||||
"--port",
|
||||
str(self.config.bridge.port),
|
||||
"--session",
|
||||
str(session_dir),
|
||||
"--mode",
|
||||
self.config.bridge.mode,
|
||||
],
|
||||
cwd=str(bridge_dir),
|
||||
env=env,
|
||||
)
|
||||
|
||||
def _resolve_bridge_script(self) -> Path:
|
||||
if self.config.bridge.script_path:
|
||||
custom_path = Path(self.config.bridge.script_path).expanduser().resolve()
|
||||
return custom_path / "bridge.js" if custom_path.is_dir() else custom_path
|
||||
|
||||
source_dir = Path(__file__).resolve().parent.parent / "bridges" / "whatsapp"
|
||||
target_dir = self.runtime_dir / "whatsapp-bridge"
|
||||
target_dir.mkdir(parents=True, exist_ok=True)
|
||||
for filename in ("bridge.js", "allowlist.js", "package.json"):
|
||||
shutil.copy2(source_dir / filename, target_dir / filename)
|
||||
return target_dir / "bridge.js"
|
||||
|
||||
def _effective_bridge_token(self) -> str:
|
||||
configured = _optional_str(self.config.bridge.token)
|
||||
if configured:
|
||||
return configured
|
||||
|
||||
token_path = self.runtime_dir / "bridge_tokens" / "whatsapp.token"
|
||||
if token_path.exists():
|
||||
token = token_path.read_text(encoding="utf-8").strip()
|
||||
if token:
|
||||
return token
|
||||
token_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
token = secrets.token_urlsafe(32)
|
||||
token_path.write_text(token, encoding="utf-8")
|
||||
try:
|
||||
token_path.chmod(0o600)
|
||||
except OSError:
|
||||
pass
|
||||
return token
|
||||
|
||||
def _session_dir(self) -> Path:
|
||||
if self.config.bridge.session_dir:
|
||||
return Path(self.config.bridge.session_dir).expanduser().resolve()
|
||||
return (self.runtime_dir / "whatsapp" / "session").resolve()
|
||||
|
||||
def _outbound_media_root(self) -> Path:
|
||||
return (self.runtime_dir / "outbound_media").resolve()
|
||||
|
||||
|
||||
def _split_text(content: str, limit: int) -> list[str]:
|
||||
text = content.strip()
|
||||
if not text:
|
||||
return [""]
|
||||
if len(text) <= limit:
|
||||
return [text]
|
||||
chunks = []
|
||||
remaining = text
|
||||
while remaining:
|
||||
chunk = remaining[:limit]
|
||||
if len(remaining) > limit:
|
||||
split_at = chunk.rfind("\n")
|
||||
if split_at < limit // 3:
|
||||
split_at = chunk.rfind(" ")
|
||||
if split_at >= limit // 3:
|
||||
chunk = chunk[:split_at]
|
||||
chunks.append(chunk.strip())
|
||||
remaining = remaining[len(chunk):].lstrip()
|
||||
return [chunk for chunk in chunks if chunk]
|
||||
|
||||
|
||||
def _optional_str(value: Any) -> Optional[str]:
|
||||
if value is None:
|
||||
return None
|
||||
value = str(value).strip()
|
||||
return value or None
|
||||
|
||||
|
||||
def _build_session_key_hint(source: ChannelSource) -> str:
|
||||
parts = [source.platform.value, source.chat_id]
|
||||
if source.thread_id:
|
||||
parts.append(source.thread_id)
|
||||
return "__".join(part.replace("/", "_") for part in parts if part)
|
||||
|
||||
|
||||
def _normalize_whatsapp_identifier(value: Any) -> str:
|
||||
normalized = re.sub(r":.*@", "@", str(value or "").strip())
|
||||
normalized = re.sub(r"@.*", "", normalized)
|
||||
return normalized.lstrip("+")
|
||||
123
openspace/communication/attachment_cache.py
Normal file
123
openspace/communication/attachment_cache.py
Normal file
|
|
@ -0,0 +1,123 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import shutil
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
from openspace.utils.logging import Logger
|
||||
|
||||
from .types import AttachmentKind, ChannelAttachment
|
||||
|
||||
logger = Logger.get_logger(__name__)
|
||||
|
||||
|
||||
class AttachmentCache:
|
||||
def __init__(
|
||||
self,
|
||||
base_dir: Path,
|
||||
*,
|
||||
max_attachment_bytes: int = 25 * 1024 * 1024,
|
||||
max_session_attachment_bytes: int = 100 * 1024 * 1024,
|
||||
):
|
||||
self.base_dir = base_dir
|
||||
self.max_attachment_bytes = max_attachment_bytes
|
||||
self.max_session_attachment_bytes = max_session_attachment_bytes
|
||||
self.base_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def session_dir(self, session_key: str) -> Path:
|
||||
directory = self.base_dir / session_key / "attachments"
|
||||
directory.mkdir(parents=True, exist_ok=True)
|
||||
return directory
|
||||
|
||||
def save_bytes(
|
||||
self,
|
||||
*,
|
||||
session_key: str,
|
||||
data: bytes,
|
||||
filename: str,
|
||||
kind: AttachmentKind,
|
||||
mime_type: str = "",
|
||||
) -> Optional[ChannelAttachment]:
|
||||
data_size = len(data)
|
||||
if not self._within_limits(session_key, data_size):
|
||||
return None
|
||||
|
||||
directory = self.session_dir(session_key)
|
||||
safe_name = _safe_name(filename)
|
||||
target = directory / f"{uuid.uuid4().hex[:12]}_{safe_name}"
|
||||
target.write_bytes(data)
|
||||
return ChannelAttachment(
|
||||
kind=kind,
|
||||
path=str(target),
|
||||
name=safe_name,
|
||||
mime_type=mime_type,
|
||||
size_bytes=len(data),
|
||||
)
|
||||
|
||||
def copy_local_file(
|
||||
self,
|
||||
*,
|
||||
session_key: str,
|
||||
source_path: str,
|
||||
kind: AttachmentKind,
|
||||
preferred_name: Optional[str] = None,
|
||||
mime_type: str = "",
|
||||
) -> Optional[ChannelAttachment]:
|
||||
source = Path(source_path).expanduser()
|
||||
if not source.exists():
|
||||
logger.warning("Attachment source does not exist: %s", source)
|
||||
return None
|
||||
source_size = source.stat().st_size
|
||||
if not self._within_limits(session_key, source_size):
|
||||
return None
|
||||
|
||||
directory = self.session_dir(session_key)
|
||||
safe_name = _safe_name(preferred_name or source.name)
|
||||
target = directory / f"{uuid.uuid4().hex[:12]}_{safe_name}"
|
||||
shutil.copy2(source, target)
|
||||
return ChannelAttachment(
|
||||
kind=kind,
|
||||
path=str(target),
|
||||
name=safe_name,
|
||||
mime_type=mime_type,
|
||||
size_bytes=target.stat().st_size,
|
||||
metadata={"source_path": str(source)},
|
||||
)
|
||||
|
||||
def _within_limits(self, session_key: str, attachment_size: int) -> bool:
|
||||
if attachment_size > self.max_attachment_bytes:
|
||||
logger.warning(
|
||||
"Rejecting attachment for session %s because %d bytes exceeds limit %d",
|
||||
session_key,
|
||||
attachment_size,
|
||||
self.max_attachment_bytes,
|
||||
)
|
||||
return False
|
||||
|
||||
session_usage = self._session_usage_bytes(session_key)
|
||||
if session_usage + attachment_size > self.max_session_attachment_bytes:
|
||||
logger.warning(
|
||||
"Rejecting attachment for session %s because session quota would exceed %d bytes",
|
||||
session_key,
|
||||
self.max_session_attachment_bytes,
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
def _session_usage_bytes(self, session_key: str) -> int:
|
||||
directory = self.base_dir / session_key / "attachments"
|
||||
if not directory.exists():
|
||||
return 0
|
||||
|
||||
total = 0
|
||||
for path in directory.iterdir():
|
||||
if path.is_file():
|
||||
total += path.stat().st_size
|
||||
return total
|
||||
|
||||
|
||||
def _safe_name(name: str) -> str:
|
||||
value = (name or "attachment").replace("\x00", "").strip()
|
||||
value = Path(value).name
|
||||
return value or "attachment"
|
||||
71
openspace/communication/bridges/whatsapp/allowlist.js
Normal file
71
openspace/communication/bridges/whatsapp/allowlist.js
Normal file
|
|
@ -0,0 +1,71 @@
|
|||
import path from 'path';
|
||||
import { existsSync, readFileSync } from 'fs';
|
||||
|
||||
export function normalizeWhatsAppIdentifier(value) {
|
||||
return String(value || '')
|
||||
.trim()
|
||||
.replace(/:.*@/, '@')
|
||||
.replace(/@.*/, '')
|
||||
.replace(/^\+/, '');
|
||||
}
|
||||
|
||||
export function parseAllowedUsers(rawValue) {
|
||||
return new Set(
|
||||
String(rawValue || '')
|
||||
.split(',')
|
||||
.map((value) => normalizeWhatsAppIdentifier(value))
|
||||
.filter(Boolean)
|
||||
);
|
||||
}
|
||||
|
||||
function readMappingFile(sessionDir, identifier, suffix = '') {
|
||||
const filePath = path.join(sessionDir, `lid-mapping-${identifier}${suffix}.json`);
|
||||
if (!existsSync(filePath)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
try {
|
||||
const parsed = JSON.parse(readFileSync(filePath, 'utf8'));
|
||||
const normalized = normalizeWhatsAppIdentifier(parsed);
|
||||
return normalized || null;
|
||||
} catch {
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
export function expandWhatsAppIdentifiers(identifier, sessionDir) {
|
||||
const normalized = normalizeWhatsAppIdentifier(identifier);
|
||||
if (!normalized) {
|
||||
return new Set();
|
||||
}
|
||||
|
||||
const resolved = new Set();
|
||||
const queue = [normalized];
|
||||
while (queue.length > 0) {
|
||||
const current = queue.shift();
|
||||
if (!current || resolved.has(current)) {
|
||||
continue;
|
||||
}
|
||||
resolved.add(current);
|
||||
for (const suffix of ['', '_reverse']) {
|
||||
const mapped = readMappingFile(sessionDir, current, suffix);
|
||||
if (mapped && !resolved.has(mapped)) {
|
||||
queue.push(mapped);
|
||||
}
|
||||
}
|
||||
}
|
||||
return resolved;
|
||||
}
|
||||
|
||||
export function matchesAllowedUser(senderId, allowedUsers, sessionDir) {
|
||||
if (!allowedUsers || allowedUsers.size === 0) {
|
||||
return true;
|
||||
}
|
||||
const aliases = expandWhatsAppIdentifiers(senderId, sessionDir);
|
||||
for (const alias of aliases) {
|
||||
if (allowedUsers.has(alias)) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
577
openspace/communication/bridges/whatsapp/bridge.js
Normal file
577
openspace/communication/bridges/whatsapp/bridge.js
Normal file
|
|
@ -0,0 +1,577 @@
|
|||
#!/usr/bin/env node
|
||||
import {
|
||||
DisconnectReason,
|
||||
downloadMediaMessage,
|
||||
fetchLatestBaileysVersion,
|
||||
makeWASocket,
|
||||
useMultiFileAuthState,
|
||||
} from '@whiskeysockets/baileys';
|
||||
import { Boom } from '@hapi/boom';
|
||||
import pino from 'pino';
|
||||
import path from 'path';
|
||||
import {
|
||||
existsSync,
|
||||
mkdirSync,
|
||||
readFileSync,
|
||||
readdirSync,
|
||||
realpathSync,
|
||||
statSync,
|
||||
writeFileSync,
|
||||
} from 'fs';
|
||||
import { randomBytes } from 'crypto';
|
||||
import qrcode from 'qrcode-terminal';
|
||||
import { WebSocketServer, WebSocket } from 'ws';
|
||||
|
||||
import {
|
||||
matchesAllowedUser,
|
||||
normalizeWhatsAppIdentifier,
|
||||
parseAllowedUsers,
|
||||
} from './allowlist.js';
|
||||
|
||||
const args = process.argv.slice(2);
|
||||
|
||||
function getArg(name, defaultValue) {
|
||||
const index = args.indexOf(`--${name}`);
|
||||
return index !== -1 && args[index + 1] ? args[index + 1] : defaultValue;
|
||||
}
|
||||
|
||||
const PORT = parseInt(getArg('port', '3000'), 10);
|
||||
const HOST = getArg('host', '127.0.0.1');
|
||||
const BIND_HOST = HOST === 'localhost' ? '127.0.0.1' : HOST;
|
||||
const SESSION_DIR = path.resolve(
|
||||
getArg('session', path.join(process.env.HOME || '~', '.openspace', 'whatsapp', 'session'))
|
||||
);
|
||||
const WHATSAPP_MODE = getArg('mode', process.env.WHATSAPP_MODE || 'self-chat');
|
||||
const BRIDGE_TOKEN = String(process.env.BRIDGE_TOKEN || '').trim();
|
||||
const DEFAULT_REPLY_PREFIX = 'OpenSpace\n────────────\n';
|
||||
const REPLY_PREFIX = process.env.WHATSAPP_REPLY_PREFIX === undefined
|
||||
? DEFAULT_REPLY_PREFIX
|
||||
: process.env.WHATSAPP_REPLY_PREFIX.replace(/\\n/g, '\n');
|
||||
const ALLOWED_USERS = parseAllowedUsers(process.env.WHATSAPP_ALLOWED_USERS || '');
|
||||
|
||||
const IMAGE_CACHE_DIR = path.join(SESSION_DIR, '..', 'image_cache');
|
||||
const DOCUMENT_CACHE_DIR = path.join(SESSION_DIR, '..', 'document_cache');
|
||||
const AUDIO_CACHE_DIR = path.join(SESSION_DIR, '..', 'audio_cache');
|
||||
|
||||
const MAX_RECENT_SENT = 50;
|
||||
const MAX_RECENT_INBOUND = 512;
|
||||
const AUTH_TIMEOUT_MS = 5000;
|
||||
|
||||
if (!BRIDGE_TOKEN) {
|
||||
console.error('BRIDGE_TOKEN is required');
|
||||
process.exit(1);
|
||||
}
|
||||
if (!['127.0.0.1', 'localhost'].includes(HOST)) {
|
||||
console.error(`Refusing to bind WhatsApp bridge to non-loopback host: ${HOST}`);
|
||||
process.exit(1);
|
||||
}
|
||||
|
||||
mkdirSync(SESSION_DIR, { recursive: true });
|
||||
mkdirSync(IMAGE_CACHE_DIR, { recursive: true });
|
||||
mkdirSync(DOCUMENT_CACHE_DIR, { recursive: true });
|
||||
mkdirSync(AUDIO_CACHE_DIR, { recursive: true });
|
||||
|
||||
const logger = pino({ level: 'warn' });
|
||||
const clients = new Set();
|
||||
const recentlySentIds = new Set();
|
||||
const recentInboundById = new Map();
|
||||
|
||||
let sock = null;
|
||||
let connectionState = 'disconnected';
|
||||
let reconnectTimer = null;
|
||||
|
||||
function broadcast(payload) {
|
||||
const encoded = JSON.stringify(payload);
|
||||
for (const ws of clients) {
|
||||
if (ws.readyState === WebSocket.OPEN && ws._authed) {
|
||||
ws.send(encoded);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
function recordRecentOutbound(messageId) {
|
||||
if (!messageId) {
|
||||
return;
|
||||
}
|
||||
recentlySentIds.delete(messageId);
|
||||
recentlySentIds.add(messageId);
|
||||
while (recentlySentIds.size > MAX_RECENT_SENT) {
|
||||
recentlySentIds.delete(recentlySentIds.values().next().value);
|
||||
}
|
||||
}
|
||||
|
||||
function recordRecentInbound(messageId, rawMessage) {
|
||||
if (!messageId || !rawMessage) {
|
||||
return;
|
||||
}
|
||||
recentInboundById.delete(messageId);
|
||||
recentInboundById.set(messageId, rawMessage);
|
||||
while (recentInboundById.size > MAX_RECENT_INBOUND) {
|
||||
const firstKey = recentInboundById.keys().next().value;
|
||||
recentInboundById.delete(firstKey);
|
||||
}
|
||||
}
|
||||
|
||||
function currentStatusPayload() {
|
||||
return { type: 'status', status: connectionState };
|
||||
}
|
||||
|
||||
function buildQuotedOptions(replyToMessageId) {
|
||||
if (!replyToMessageId) {
|
||||
return {};
|
||||
}
|
||||
const quoted = recentInboundById.get(String(replyToMessageId).trim());
|
||||
return quoted ? { quoted } : {};
|
||||
}
|
||||
|
||||
function formatOutgoingMessage(message) {
|
||||
if (WHATSAPP_MODE !== 'self-chat') {
|
||||
return message;
|
||||
}
|
||||
return REPLY_PREFIX ? `${REPLY_PREFIX}${message}` : message;
|
||||
}
|
||||
|
||||
function buildLidMap() {
|
||||
const mapping = {};
|
||||
try {
|
||||
for (const fileName of readdirSync(SESSION_DIR)) {
|
||||
const match = fileName.match(/^lid-mapping-(.+)\.json$/);
|
||||
if (!match) {
|
||||
continue;
|
||||
}
|
||||
const value = JSON.parse(readFileSync(path.join(SESSION_DIR, fileName), 'utf8'));
|
||||
const normalized = normalizeWhatsAppIdentifier(value);
|
||||
if (normalized) {
|
||||
mapping[normalized] = match[1];
|
||||
mapping[match[1]] = normalized;
|
||||
}
|
||||
}
|
||||
} catch {
|
||||
return {};
|
||||
}
|
||||
return mapping;
|
||||
}
|
||||
|
||||
let lidToPhone = buildLidMap();
|
||||
|
||||
function normalizeId(value) {
|
||||
return normalizeWhatsAppIdentifier(value);
|
||||
}
|
||||
|
||||
function getMyIdentifiers() {
|
||||
const ids = new Set();
|
||||
if (sock?.user?.id) {
|
||||
ids.add(normalizeId(sock.user.id));
|
||||
}
|
||||
if (sock?.user?.lid) {
|
||||
ids.add(normalizeId(sock.user.lid));
|
||||
}
|
||||
return ids;
|
||||
}
|
||||
|
||||
function getMessageContainer(message) {
|
||||
return message?.message || {};
|
||||
}
|
||||
|
||||
function extractContextInfo(message) {
|
||||
const container = getMessageContainer(message);
|
||||
return (
|
||||
container.extendedTextMessage?.contextInfo
|
||||
|| container.imageMessage?.contextInfo
|
||||
|| container.videoMessage?.contextInfo
|
||||
|| container.documentMessage?.contextInfo
|
||||
|| container.audioMessage?.contextInfo
|
||||
|| container.conversation?.contextInfo
|
||||
|| null
|
||||
);
|
||||
}
|
||||
|
||||
async function cacheMedia(rawMessage, mediaMessage, targetDir, prefix, fallbackExt) {
|
||||
const buffer = await downloadMediaMessage(
|
||||
rawMessage,
|
||||
'buffer',
|
||||
{},
|
||||
{ logger, reuploadRequest: sock.updateMediaMessage }
|
||||
);
|
||||
mkdirSync(targetDir, { recursive: true });
|
||||
const mime = mediaMessage.mimetype || '';
|
||||
const ext = mime.includes('/') ? `.${mime.split('/')[1].split(';')[0]}` : fallbackExt;
|
||||
const filePath = path.join(targetDir, `${prefix}_${randomBytes(6).toString('hex')}${ext || fallbackExt}`);
|
||||
writeFileSync(filePath, buffer);
|
||||
return filePath;
|
||||
}
|
||||
|
||||
function resolveAllowedFilePath(filePath) {
|
||||
const mediaRootRaw = String(process.env.BRIDGE_MEDIA_ROOT || '').trim();
|
||||
if (!mediaRootRaw) {
|
||||
throw new Error('BRIDGE_MEDIA_ROOT is not configured');
|
||||
}
|
||||
const mediaRoot = realpathSync(mediaRootRaw);
|
||||
const resolvedPath = realpathSync(String(filePath || ''));
|
||||
const relative = path.relative(mediaRoot, resolvedPath);
|
||||
if (!relative || relative.startsWith('..') || path.isAbsolute(relative)) {
|
||||
throw new Error(`File path escapes bridge media root: ${filePath}`);
|
||||
}
|
||||
const stat = statSync(resolvedPath);
|
||||
if (!stat.isFile()) {
|
||||
throw new Error(`File is not a regular file: ${filePath}`);
|
||||
}
|
||||
return resolvedPath;
|
||||
}
|
||||
|
||||
async function sendText(to, text, replyToMessageId) {
|
||||
if (!sock || connectionState !== 'connected') {
|
||||
throw new Error('Not connected to WhatsApp');
|
||||
}
|
||||
const sent = await sock.sendMessage(
|
||||
to,
|
||||
{ text: formatOutgoingMessage(text) },
|
||||
buildQuotedOptions(replyToMessageId)
|
||||
);
|
||||
recordRecentOutbound(sent?.key?.id);
|
||||
return sent;
|
||||
}
|
||||
|
||||
async function sendMedia(to, filePath, mimetype, caption, fileName, replyToMessageId) {
|
||||
if (!sock || connectionState !== 'connected') {
|
||||
throw new Error('Not connected to WhatsApp');
|
||||
}
|
||||
|
||||
const resolvedPath = resolveAllowedFilePath(filePath);
|
||||
if (!existsSync(resolvedPath)) {
|
||||
throw new Error(`File not found: ${resolvedPath}`);
|
||||
}
|
||||
const buffer = readFileSync(resolvedPath);
|
||||
const normalizedMime = String(mimetype || '').toLowerCase();
|
||||
let payload;
|
||||
if (normalizedMime.startsWith('image/')) {
|
||||
payload = { image: buffer, caption: caption || undefined };
|
||||
} else if (normalizedMime.startsWith('video/')) {
|
||||
payload = { video: buffer, caption: caption || undefined };
|
||||
} else if (normalizedMime.startsWith('audio/')) {
|
||||
payload = {
|
||||
audio: buffer,
|
||||
mimetype: normalizedMime || 'audio/ogg; codecs=opus',
|
||||
ptt: normalizedMime.includes('ogg') || normalizedMime.includes('opus'),
|
||||
};
|
||||
} else {
|
||||
payload = {
|
||||
document: buffer,
|
||||
fileName: fileName || path.basename(resolvedPath),
|
||||
caption: caption || undefined,
|
||||
};
|
||||
}
|
||||
|
||||
const sent = await sock.sendMessage(
|
||||
to,
|
||||
payload,
|
||||
buildQuotedOptions(replyToMessageId)
|
||||
);
|
||||
recordRecentOutbound(sent?.key?.id);
|
||||
return sent;
|
||||
}
|
||||
|
||||
async function handleInboundMessage(rawMessage) {
|
||||
if (!rawMessage?.message) {
|
||||
return;
|
||||
}
|
||||
|
||||
const chatId = rawMessage.key?.remoteJid || '';
|
||||
const senderId = rawMessage.key?.participant || chatId;
|
||||
const isGroup = chatId.endsWith('@g.us');
|
||||
const senderNumber = senderId.replace(/@.*/, '');
|
||||
|
||||
if (rawMessage.key?.fromMe) {
|
||||
if (isGroup || chatId.includes('status')) {
|
||||
return;
|
||||
}
|
||||
if (WHATSAPP_MODE === 'bot') {
|
||||
return;
|
||||
}
|
||||
|
||||
const myIds = getMyIdentifiers();
|
||||
const chatNumber = normalizeId(chatId);
|
||||
if (!myIds.has(chatNumber)) {
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
if (!rawMessage.key?.fromMe && !matchesAllowedUser(senderId, ALLOWED_USERS, SESSION_DIR)) {
|
||||
return;
|
||||
}
|
||||
|
||||
const container = getMessageContainer(rawMessage);
|
||||
const contextInfo = extractContextInfo(rawMessage);
|
||||
const mentionedIds = (contextInfo?.mentionedJid || []).map((value) => normalizeId(value));
|
||||
const mentionsBot = mentionedIds.some((value) => getMyIdentifiers().has(value));
|
||||
const replyToMessageId = contextInfo?.stanzaId || null;
|
||||
|
||||
let body = '';
|
||||
let hasMedia = false;
|
||||
let mediaType = '';
|
||||
const mediaUrls = [];
|
||||
|
||||
if (container.conversation) {
|
||||
body = container.conversation;
|
||||
} else if (container.extendedTextMessage?.text) {
|
||||
body = container.extendedTextMessage.text;
|
||||
} else if (container.imageMessage) {
|
||||
body = container.imageMessage.caption || '';
|
||||
hasMedia = true;
|
||||
mediaType = 'image';
|
||||
try {
|
||||
mediaUrls.push(await cacheMedia(rawMessage, container.imageMessage, IMAGE_CACHE_DIR, 'img', '.jpg'));
|
||||
} catch (error) {
|
||||
console.error('[bridge] Failed to download image:', error.message);
|
||||
}
|
||||
} else if (container.videoMessage) {
|
||||
body = container.videoMessage.caption || '';
|
||||
hasMedia = true;
|
||||
mediaType = 'video';
|
||||
try {
|
||||
mediaUrls.push(await cacheMedia(rawMessage, container.videoMessage, DOCUMENT_CACHE_DIR, 'vid', '.mp4'));
|
||||
} catch (error) {
|
||||
console.error('[bridge] Failed to download video:', error.message);
|
||||
}
|
||||
} else if (container.audioMessage || container.pttMessage) {
|
||||
hasMedia = true;
|
||||
mediaType = container.pttMessage ? 'ptt' : 'audio';
|
||||
try {
|
||||
const audioMessage = container.pttMessage || container.audioMessage;
|
||||
mediaUrls.push(await cacheMedia(rawMessage, audioMessage, AUDIO_CACHE_DIR, 'aud', '.ogg'));
|
||||
} catch (error) {
|
||||
console.error('[bridge] Failed to download audio:', error.message);
|
||||
}
|
||||
} else if (container.documentMessage) {
|
||||
body = container.documentMessage.caption || '';
|
||||
hasMedia = true;
|
||||
mediaType = 'document';
|
||||
try {
|
||||
mediaUrls.push(
|
||||
await cacheMedia(
|
||||
rawMessage,
|
||||
container.documentMessage,
|
||||
DOCUMENT_CACHE_DIR,
|
||||
'doc',
|
||||
path.extname(container.documentMessage.fileName || '') || '.bin'
|
||||
)
|
||||
);
|
||||
} catch (error) {
|
||||
console.error('[bridge] Failed to download document:', error.message);
|
||||
}
|
||||
}
|
||||
|
||||
const messageId = rawMessage.key?.id;
|
||||
if (messageId && recentlySentIds.has(messageId)) {
|
||||
recentlySentIds.delete(messageId);
|
||||
return;
|
||||
}
|
||||
|
||||
if (messageId) {
|
||||
recordRecentInbound(messageId, rawMessage);
|
||||
}
|
||||
|
||||
const normalizedSenderId = lidToPhone[normalizeId(senderId)] || senderId;
|
||||
const event = {
|
||||
type: 'message',
|
||||
messageId,
|
||||
chatId,
|
||||
senderId: normalizedSenderId,
|
||||
senderName: rawMessage.pushName || senderNumber,
|
||||
chatName: isGroup ? chatId.split('@')[0] : (rawMessage.pushName || senderNumber),
|
||||
isGroup,
|
||||
body,
|
||||
hasMedia,
|
||||
mediaType,
|
||||
mediaUrls,
|
||||
replyToMessageId,
|
||||
mentionedIds,
|
||||
mentionsBot,
|
||||
timestamp: rawMessage.messageTimestamp,
|
||||
};
|
||||
broadcast(event);
|
||||
}
|
||||
|
||||
async function startSocket() {
|
||||
const { state, saveCreds } = await useMultiFileAuthState(SESSION_DIR);
|
||||
const { version } = await fetchLatestBaileysVersion();
|
||||
|
||||
sock = makeWASocket({
|
||||
version,
|
||||
auth: state,
|
||||
logger,
|
||||
printQRInTerminal: false,
|
||||
browser: ['OpenSpace', 'Chrome', '120.0'],
|
||||
syncFullHistory: false,
|
||||
markOnlineOnConnect: false,
|
||||
getMessage: async () => ({ conversation: '' }),
|
||||
});
|
||||
|
||||
sock.ev.on('creds.update', () => {
|
||||
saveCreds();
|
||||
lidToPhone = buildLidMap();
|
||||
});
|
||||
|
||||
sock.ev.on('connection.update', (update) => {
|
||||
const { connection, lastDisconnect, qr } = update;
|
||||
if (qr) {
|
||||
console.log('\nScan this QR code with WhatsApp on your phone:\n');
|
||||
qrcode.generate(qr, { small: true });
|
||||
console.log('\nWaiting for scan...\n');
|
||||
broadcast({ type: 'qr', qr });
|
||||
}
|
||||
|
||||
if (connection === 'close') {
|
||||
const reason = new Boom(lastDisconnect?.error)?.output?.statusCode;
|
||||
connectionState = 'disconnected';
|
||||
broadcast(currentStatusPayload());
|
||||
if (reason === DisconnectReason.loggedOut) {
|
||||
console.log('Logged out. Delete session and restart to re-authenticate.');
|
||||
process.exit(1);
|
||||
}
|
||||
clearTimeout(reconnectTimer);
|
||||
reconnectTimer = setTimeout(
|
||||
() => startSocket().catch((error) => console.error('WhatsApp reconnect failed:', error)),
|
||||
reason === 515 ? 1000 : 3000
|
||||
);
|
||||
} else if (connection === 'open') {
|
||||
connectionState = 'connected';
|
||||
console.log('WhatsApp connected');
|
||||
broadcast(currentStatusPayload());
|
||||
}
|
||||
});
|
||||
|
||||
sock.ev.on('messages.upsert', async ({ messages, type }) => {
|
||||
if (type !== 'notify' && type !== 'append') {
|
||||
return;
|
||||
}
|
||||
for (const message of messages) {
|
||||
try {
|
||||
await handleInboundMessage(message);
|
||||
} catch (error) {
|
||||
console.error('[bridge] Failed to normalize inbound message:', error);
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
function startBridgeServer() {
|
||||
const wss = new WebSocketServer({ host: BIND_HOST, port: PORT });
|
||||
console.log(`OpenSpace WhatsApp bridge listening on ws://${BIND_HOST}:${PORT} (mode: ${WHATSAPP_MODE})`);
|
||||
|
||||
wss.on('connection', (ws, request) => {
|
||||
if (request.headers.origin) {
|
||||
ws.close(4003, 'Origin header is not allowed');
|
||||
return;
|
||||
}
|
||||
|
||||
ws._authed = false;
|
||||
clients.add(ws);
|
||||
const authTimeout = setTimeout(() => {
|
||||
if (!ws._authed) {
|
||||
ws.close(4001, 'Authentication timeout');
|
||||
}
|
||||
}, AUTH_TIMEOUT_MS);
|
||||
|
||||
ws.once('message', (raw) => {
|
||||
try {
|
||||
const payload = JSON.parse(raw.toString());
|
||||
if (payload.type !== 'auth' || payload.token !== BRIDGE_TOKEN) {
|
||||
ws.close(4003, 'Invalid bridge token');
|
||||
return;
|
||||
}
|
||||
ws._authed = true;
|
||||
clearTimeout(authTimeout);
|
||||
ws.send(JSON.stringify({ type: 'auth_ok' }));
|
||||
ws.send(JSON.stringify(currentStatusPayload()));
|
||||
|
||||
ws.on('message', async (commandRaw) => {
|
||||
try {
|
||||
const command = JSON.parse(commandRaw.toString());
|
||||
const requestId = command.requestId || null;
|
||||
let sent = null;
|
||||
if (command.type === 'send') {
|
||||
sent = await sendText(command.to, command.text || '', command.replyToMessageId);
|
||||
} else if (command.type === 'send_media') {
|
||||
sent = await sendMedia(
|
||||
command.to,
|
||||
command.filePath,
|
||||
command.mimetype,
|
||||
command.caption,
|
||||
command.fileName,
|
||||
command.replyToMessageId
|
||||
);
|
||||
} else {
|
||||
throw new Error(`Unsupported bridge command: ${command.type}`);
|
||||
}
|
||||
ws.send(JSON.stringify({
|
||||
type: 'ack',
|
||||
requestId,
|
||||
messageId: sent?.key?.id || null,
|
||||
}));
|
||||
} catch (error) {
|
||||
ws.send(JSON.stringify({
|
||||
type: 'error',
|
||||
requestId: (() => {
|
||||
try {
|
||||
return JSON.parse(commandRaw.toString()).requestId || null;
|
||||
} catch {
|
||||
return null;
|
||||
}
|
||||
})(),
|
||||
error: error?.message || String(error),
|
||||
}));
|
||||
}
|
||||
});
|
||||
} catch {
|
||||
ws.close(4003, 'Invalid auth payload');
|
||||
}
|
||||
});
|
||||
|
||||
ws.on('close', () => {
|
||||
clearTimeout(authTimeout);
|
||||
clients.delete(ws);
|
||||
});
|
||||
ws.on('error', () => {
|
||||
clearTimeout(authTimeout);
|
||||
clients.delete(ws);
|
||||
});
|
||||
});
|
||||
|
||||
return wss;
|
||||
}
|
||||
|
||||
async function shutdown(server) {
|
||||
clearTimeout(reconnectTimer);
|
||||
for (const ws of clients) {
|
||||
try {
|
||||
ws.close();
|
||||
} catch {}
|
||||
}
|
||||
clients.clear();
|
||||
if (server) {
|
||||
await new Promise((resolve) => server.close(resolve));
|
||||
}
|
||||
if (sock) {
|
||||
try {
|
||||
sock.end(new Error('Bridge shutdown'));
|
||||
} catch {}
|
||||
sock = null;
|
||||
}
|
||||
}
|
||||
|
||||
const server = startBridgeServer();
|
||||
startSocket().catch((error) => {
|
||||
console.error('Failed to start WhatsApp socket:', error);
|
||||
process.exit(1);
|
||||
});
|
||||
|
||||
for (const signal of ['SIGINT', 'SIGTERM']) {
|
||||
process.on(signal, async () => {
|
||||
try {
|
||||
await shutdown(server);
|
||||
} finally {
|
||||
process.exit(0);
|
||||
}
|
||||
});
|
||||
}
|
||||
17
openspace/communication/bridges/whatsapp/package.json
Normal file
17
openspace/communication/bridges/whatsapp/package.json
Normal file
|
|
@ -0,0 +1,17 @@
|
|||
{
|
||||
"name": "openspace-whatsapp-bridge",
|
||||
"version": "1.0.0",
|
||||
"description": "WhatsApp bridge for OpenSpace using Baileys",
|
||||
"private": true,
|
||||
"type": "module",
|
||||
"scripts": {
|
||||
"start": "node bridge.js"
|
||||
},
|
||||
"dependencies": {
|
||||
"@hapi/boom": "^10.0.1",
|
||||
"@whiskeysockets/baileys": "7.0.0-rc.9",
|
||||
"pino": "^9.0.0",
|
||||
"qrcode-terminal": "^0.12.0",
|
||||
"ws": "^8.18.0"
|
||||
}
|
||||
}
|
||||
375
openspace/communication/config.py
Normal file
375
openspace/communication/config.py
Normal file
|
|
@ -0,0 +1,375 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Literal, Optional
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator, model_validator
|
||||
|
||||
from openspace.host_detection import load_runtime_env
|
||||
from openspace.utils.logging import Logger
|
||||
|
||||
logger = Logger.get_logger(__name__)
|
||||
|
||||
|
||||
class GatewayServerConfig(BaseModel):
|
||||
host: str = "127.0.0.1"
|
||||
port: int = Field(8765, ge=1, le=65535)
|
||||
health_path: str = "/health"
|
||||
|
||||
@field_validator("health_path")
|
||||
@classmethod
|
||||
def validate_health_path(cls, value: str) -> str:
|
||||
value = value.strip() or "/health"
|
||||
if not value.startswith("/"):
|
||||
value = "/" + value
|
||||
return value
|
||||
|
||||
|
||||
class AgentExecutionConfig(BaseModel):
|
||||
max_iterations: int = Field(20, ge=1, le=200)
|
||||
enable_recording: bool = True
|
||||
recording_backends: List[str] = Field(default_factory=lambda: ["shell"])
|
||||
backend_scope: Optional[List[str]] = None
|
||||
grounding_config_path: Optional[str] = None
|
||||
workspace_root: Optional[str] = None
|
||||
llm_timeout: float = Field(120.0, ge=1.0, le=3600.0)
|
||||
|
||||
|
||||
class SessionProcessingConfig(BaseModel):
|
||||
history_max_turns: int = Field(12, ge=1, le=100)
|
||||
max_parallel_sessions: int = Field(2, ge=1, le=64)
|
||||
idle_ttl_seconds: int = Field(900, ge=30, le=86400)
|
||||
per_session_queue_size: int = Field(32, ge=1, le=512)
|
||||
whatsapp_poll_interval_seconds: float = Field(1.0, ge=0.1, le=60.0)
|
||||
max_attachment_bytes: int = Field(25 * 1024 * 1024, ge=1, le=512 * 1024 * 1024)
|
||||
max_session_attachment_bytes: int = Field(
|
||||
100 * 1024 * 1024,
|
||||
ge=1,
|
||||
le=10 * 1024 * 1024 * 1024,
|
||||
)
|
||||
|
||||
|
||||
class ChannelAccessConfig(BaseModel):
|
||||
enabled: bool = False
|
||||
allow_all_users: bool = False
|
||||
allowed_users: List[str] = Field(default_factory=list)
|
||||
allow_dm: bool = True
|
||||
allow_groups: bool = True
|
||||
group_policy: Literal["disabled", "mention_only", "reply_or_mention", "all"] = "reply_or_mention"
|
||||
|
||||
|
||||
class WhatsAppBridgeConfig(BaseModel):
|
||||
host: str = "127.0.0.1"
|
||||
port: int = Field(3000, ge=1, le=65535)
|
||||
script_path: Optional[str] = None
|
||||
session_dir: Optional[str] = None
|
||||
mode: Literal["self-chat", "bot"] = "self-chat"
|
||||
auto_install_dependencies: bool = True
|
||||
token: Optional[str] = None
|
||||
enforce_loopback: bool = True
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_loopback_constraints(self) -> "WhatsAppBridgeConfig":
|
||||
host = self.host.strip().lower() or "127.0.0.1"
|
||||
if self.enforce_loopback and host not in {"127.0.0.1", "localhost"}:
|
||||
raise ValueError(
|
||||
"WhatsApp bridge host must be loopback when enforce_loopback is enabled"
|
||||
)
|
||||
self.host = host
|
||||
return self
|
||||
|
||||
@property
|
||||
def base_url(self) -> str:
|
||||
return f"http://{self.host}:{self.port}"
|
||||
|
||||
@property
|
||||
def ws_url(self) -> str:
|
||||
return f"ws://{self.listen_host}:{self.port}"
|
||||
|
||||
@property
|
||||
def listen_host(self) -> str:
|
||||
return "127.0.0.1" if self.host == "localhost" else self.host
|
||||
|
||||
|
||||
class WhatsAppConfig(ChannelAccessConfig):
|
||||
bridge: WhatsAppBridgeConfig = Field(default_factory=WhatsAppBridgeConfig)
|
||||
reply_prefix: Optional[str] = None
|
||||
|
||||
|
||||
class FeishuConfig(ChannelAccessConfig):
|
||||
app_id: Optional[str] = None
|
||||
app_secret: Optional[str] = None
|
||||
domain: Literal["feishu", "lark"] = "feishu"
|
||||
connection_mode: Literal["webhook", "websocket"] = "webhook"
|
||||
verification_token: Optional[str] = None
|
||||
encrypt_key: Optional[str] = None
|
||||
bot_open_id: Optional[str] = None
|
||||
webhook_path: str = "/feishu/webhook"
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_webhook_requirements(self) -> "FeishuConfig":
|
||||
if self.enabled and self.connection_mode == "webhook" and not (self.verification_token or "").strip():
|
||||
raise ValueError("Feishu webhook mode requires verification_token")
|
||||
return self
|
||||
|
||||
@field_validator("webhook_path")
|
||||
@classmethod
|
||||
def validate_webhook_path(cls, value: str) -> str:
|
||||
value = value.strip() or "/feishu/webhook"
|
||||
if not value.startswith("/"):
|
||||
value = "/" + value
|
||||
return value
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_webhook_security(self) -> "FeishuConfig":
|
||||
if self.enabled and self.connection_mode == "webhook":
|
||||
token = (self.verification_token or "").strip()
|
||||
if not token:
|
||||
raise ValueError(
|
||||
"Feishu webhook mode requires verification_token when enabled"
|
||||
)
|
||||
self.verification_token = token
|
||||
if self.encrypt_key is not None:
|
||||
self.encrypt_key = self.encrypt_key.strip() or None
|
||||
if self.bot_open_id is not None:
|
||||
self.bot_open_id = self.bot_open_id.strip() or None
|
||||
return self
|
||||
|
||||
|
||||
class CommunicationConfig(BaseModel):
|
||||
data_dir: str = Field(
|
||||
default_factory=lambda: str(
|
||||
Path(__file__).resolve().parents[2] / "logs" / "communication"
|
||||
)
|
||||
)
|
||||
server: GatewayServerConfig = Field(default_factory=GatewayServerConfig)
|
||||
agent: AgentExecutionConfig = Field(default_factory=AgentExecutionConfig)
|
||||
sessions: SessionProcessingConfig = Field(default_factory=SessionProcessingConfig)
|
||||
whatsapp: WhatsAppConfig = Field(default_factory=WhatsAppConfig)
|
||||
feishu: FeishuConfig = Field(default_factory=FeishuConfig)
|
||||
|
||||
@property
|
||||
def openspace(self) -> AgentExecutionConfig:
|
||||
return self.agent
|
||||
|
||||
@property
|
||||
def runtime(self) -> SessionProcessingConfig:
|
||||
return self.sessions
|
||||
|
||||
@property
|
||||
def data_path(self) -> Path:
|
||||
return Path(self.data_dir).expanduser().resolve()
|
||||
|
||||
@property
|
||||
def sessions_dir(self) -> Path:
|
||||
return self.data_path / "sessions"
|
||||
|
||||
@property
|
||||
def bridge_assets_dir(self) -> Path:
|
||||
return Path(__file__).resolve().parent / "bridges" / "whatsapp"
|
||||
|
||||
@property
|
||||
def runtime_status_path(self) -> Path:
|
||||
return self.data_path / "runtime_status.json"
|
||||
|
||||
@property
|
||||
def locks_dir(self) -> Path:
|
||||
return self.data_path / "locks"
|
||||
|
||||
@property
|
||||
def bridge_tokens_dir(self) -> Path:
|
||||
return self.data_path / "bridge_tokens"
|
||||
|
||||
@property
|
||||
def whatsapp_bridge_token_path(self) -> Path:
|
||||
return self.bridge_tokens_dir / "whatsapp.token"
|
||||
|
||||
@property
|
||||
def outbound_media_dir(self) -> Path:
|
||||
return self.data_path / "outbound_media"
|
||||
|
||||
@property
|
||||
def feishu_seen_message_ids_path(self) -> Path:
|
||||
return self.data_path / "feishu_seen_message_ids.json"
|
||||
|
||||
@property
|
||||
def enabled_platforms(self) -> List[str]:
|
||||
platforms: List[str] = []
|
||||
if self.whatsapp.enabled:
|
||||
platforms.append("whatsapp")
|
||||
if self.feishu.enabled:
|
||||
platforms.append("feishu")
|
||||
return platforms
|
||||
|
||||
|
||||
def load_communication_config(path: Optional[str] = None) -> CommunicationConfig:
|
||||
load_runtime_env()
|
||||
config_path = _resolve_config_path(path)
|
||||
raw: Dict[str, Any] = {}
|
||||
|
||||
if config_path and config_path.is_file():
|
||||
with open(config_path, "r", encoding="utf-8") as handle:
|
||||
raw = json.load(handle) or {}
|
||||
raw = _normalize_legacy_keys(raw)
|
||||
logger.info("Loaded communication config: %s", config_path)
|
||||
|
||||
config = CommunicationConfig.model_validate(raw)
|
||||
_apply_env_overrides(config)
|
||||
return CommunicationConfig.model_validate(config.model_dump(mode="python"))
|
||||
|
||||
|
||||
def _resolve_config_path(path: Optional[str]) -> Optional[Path]:
|
||||
explicit_path = Path(path).expanduser() if path else None
|
||||
if explicit_path is not None:
|
||||
if not explicit_path.is_file():
|
||||
raise FileNotFoundError(f"Communication config file not found: {explicit_path}")
|
||||
return explicit_path
|
||||
|
||||
env_path = (
|
||||
Path(os.environ["OPENSPACE_COMMUNICATION_CONFIG"]).expanduser()
|
||||
if os.environ.get("OPENSPACE_COMMUNICATION_CONFIG")
|
||||
else None
|
||||
)
|
||||
if env_path is not None:
|
||||
if not env_path.is_file():
|
||||
raise FileNotFoundError(f"Communication config file not found: {env_path}")
|
||||
return env_path
|
||||
|
||||
default_path = Path(__file__).resolve().parents[1] / "config" / "config_communication.json"
|
||||
return default_path if default_path.is_file() else None
|
||||
|
||||
|
||||
def _apply_env_overrides(config: CommunicationConfig) -> None:
|
||||
_maybe_set_bool(config.whatsapp, "enabled", os.getenv("WHATSAPP_ENABLED"))
|
||||
_maybe_set_bool(config.whatsapp, "allow_all_users", os.getenv("WHATSAPP_ALLOW_ALL_USERS"))
|
||||
_maybe_set_list(config.whatsapp, "allowed_users", os.getenv("WHATSAPP_ALLOWED_USERS"))
|
||||
_maybe_set_bool(config.whatsapp, "allow_dm", os.getenv("WHATSAPP_ALLOW_DM"))
|
||||
_maybe_set_bool(config.whatsapp, "allow_groups", os.getenv("WHATSAPP_ALLOW_GROUPS"))
|
||||
_maybe_set_str(config.whatsapp, "group_policy", os.getenv("WHATSAPP_GROUP_POLICY"))
|
||||
_maybe_set_str(config.whatsapp.bridge, "host", os.getenv("WHATSAPP_BRIDGE_HOST"))
|
||||
_maybe_set_int(config.whatsapp.bridge, "port", os.getenv("WHATSAPP_BRIDGE_PORT"))
|
||||
_maybe_set_str(config.whatsapp.bridge, "script_path", os.getenv("WHATSAPP_BRIDGE_SCRIPT"))
|
||||
_maybe_set_str(config.whatsapp.bridge, "session_dir", os.getenv("WHATSAPP_SESSION_DIR"))
|
||||
_maybe_set_str(config.whatsapp.bridge, "mode", os.getenv("WHATSAPP_MODE"))
|
||||
_maybe_set_str(config.whatsapp.bridge, "token", os.getenv("WHATSAPP_BRIDGE_TOKEN"))
|
||||
_maybe_set_bool(config.whatsapp.bridge, "enforce_loopback", os.getenv("WHATSAPP_BRIDGE_ENFORCE_LOOPBACK"))
|
||||
_maybe_set_str(config.whatsapp, "reply_prefix", os.getenv("WHATSAPP_REPLY_PREFIX"))
|
||||
|
||||
_maybe_set_bool(config.feishu, "enabled", os.getenv("FEISHU_ENABLED"))
|
||||
_maybe_set_bool(config.feishu, "allow_all_users", os.getenv("FEISHU_ALLOW_ALL_USERS"))
|
||||
_maybe_set_list(config.feishu, "allowed_users", os.getenv("FEISHU_ALLOWED_USERS"))
|
||||
_maybe_set_bool(config.feishu, "allow_dm", os.getenv("FEISHU_ALLOW_DM"))
|
||||
_maybe_set_bool(config.feishu, "allow_groups", os.getenv("FEISHU_ALLOW_GROUPS"))
|
||||
_maybe_set_str(config.feishu, "group_policy", os.getenv("FEISHU_GROUP_POLICY"))
|
||||
_maybe_set_str(config.feishu, "app_id", os.getenv("FEISHU_APP_ID"))
|
||||
_maybe_set_str(config.feishu, "app_secret", os.getenv("FEISHU_APP_SECRET"))
|
||||
_maybe_set_str(config.feishu, "verification_token", os.getenv("FEISHU_VERIFICATION_TOKEN"))
|
||||
_maybe_set_str(config.feishu, "encrypt_key", os.getenv("FEISHU_ENCRYPT_KEY"))
|
||||
_maybe_set_str(config.feishu, "bot_open_id", os.getenv("FEISHU_BOT_OPEN_ID"))
|
||||
_maybe_set_str(config.feishu, "domain", os.getenv("FEISHU_DOMAIN"))
|
||||
_maybe_set_str(config.feishu, "connection_mode", os.getenv("FEISHU_CONNECTION_MODE"))
|
||||
_maybe_set_str(config.feishu, "webhook_path", os.getenv("FEISHU_WEBHOOK_PATH"))
|
||||
|
||||
_maybe_set_str(config, "data_dir", os.getenv("OPENSPACE_COMMUNICATION_DATA_DIR"))
|
||||
_maybe_set_str(config.server, "host", os.getenv("OPENSPACE_COMMUNICATION_HOST"))
|
||||
_maybe_set_int(config.server, "port", os.getenv("OPENSPACE_COMMUNICATION_PORT"))
|
||||
_maybe_set_int(
|
||||
config.agent,
|
||||
"max_iterations",
|
||||
os.getenv("OPENSPACE_COMMUNICATION_MAX_ITERATIONS") or os.getenv("OPENSPACE_MAX_ITERATIONS"),
|
||||
)
|
||||
_maybe_set_bool(
|
||||
config.agent,
|
||||
"enable_recording",
|
||||
os.getenv("OPENSPACE_COMMUNICATION_ENABLE_RECORDING") or os.getenv("OPENSPACE_ENABLE_RECORDING"),
|
||||
)
|
||||
_maybe_set_list(
|
||||
config.agent,
|
||||
"recording_backends",
|
||||
os.getenv("OPENSPACE_COMMUNICATION_RECORDING_BACKENDS"),
|
||||
)
|
||||
_maybe_set_list(
|
||||
config.agent,
|
||||
"backend_scope",
|
||||
os.getenv("OPENSPACE_COMMUNICATION_BACKEND_SCOPE") or os.getenv("OPENSPACE_BACKEND_SCOPE"),
|
||||
)
|
||||
_maybe_set_str(
|
||||
config.agent,
|
||||
"grounding_config_path",
|
||||
os.getenv("OPENSPACE_COMMUNICATION_GROUNDING_CONFIG_PATH") or os.getenv("OPENSPACE_CONFIG_PATH"),
|
||||
)
|
||||
_maybe_set_str(config.agent, "workspace_root", os.getenv("OPENSPACE_COMMUNICATION_WORKSPACE_ROOT"))
|
||||
_maybe_set_float(config.agent, "llm_timeout", os.getenv("OPENSPACE_COMMUNICATION_LLM_TIMEOUT"))
|
||||
_maybe_set_int(config.sessions, "history_max_turns", os.getenv("OPENSPACE_COMMUNICATION_HISTORY_TURNS"))
|
||||
_maybe_set_int(config.sessions, "max_parallel_sessions", os.getenv("OPENSPACE_COMMUNICATION_MAX_PARALLEL"))
|
||||
_maybe_set_int(config.sessions, "idle_ttl_seconds", os.getenv("OPENSPACE_COMMUNICATION_IDLE_TTL"))
|
||||
_maybe_set_int(config.sessions, "per_session_queue_size", os.getenv("OPENSPACE_COMMUNICATION_QUEUE_SIZE"))
|
||||
_maybe_set_int(
|
||||
config.sessions,
|
||||
"max_attachment_bytes",
|
||||
os.getenv("OPENSPACE_COMMUNICATION_MAX_ATTACHMENT_BYTES"),
|
||||
)
|
||||
_maybe_set_int(
|
||||
config.sessions,
|
||||
"max_session_attachment_bytes",
|
||||
os.getenv("OPENSPACE_COMMUNICATION_MAX_SESSION_ATTACHMENT_BYTES"),
|
||||
)
|
||||
_maybe_set_float(
|
||||
config.sessions,
|
||||
"whatsapp_poll_interval_seconds",
|
||||
os.getenv("OPENSPACE_COMMUNICATION_WHATSAPP_POLL_INTERVAL"),
|
||||
)
|
||||
|
||||
|
||||
def _normalize_legacy_keys(raw: Dict[str, Any]) -> Dict[str, Any]:
|
||||
normalized = dict(raw)
|
||||
if "agent" not in normalized and "openspace" in normalized:
|
||||
normalized["agent"] = normalized["openspace"]
|
||||
if "sessions" not in normalized and "runtime" in normalized:
|
||||
normalized["sessions"] = normalized["runtime"]
|
||||
return normalized
|
||||
|
||||
|
||||
def _maybe_set_bool(target: Any, field_name: str, raw: Optional[str]) -> None:
|
||||
if raw is None:
|
||||
return
|
||||
lowered = raw.strip().lower()
|
||||
if lowered in {"true", "1", "yes", "on"}:
|
||||
setattr(target, field_name, True)
|
||||
elif lowered in {"false", "0", "no", "off"}:
|
||||
setattr(target, field_name, False)
|
||||
|
||||
|
||||
def _maybe_set_int(target: Any, field_name: str, raw: Optional[str]) -> None:
|
||||
if raw is None or not raw.strip():
|
||||
return
|
||||
try:
|
||||
setattr(target, field_name, int(raw))
|
||||
except ValueError:
|
||||
logger.warning("Invalid integer for %s: %r", field_name, raw)
|
||||
|
||||
|
||||
def _maybe_set_list(target: Any, field_name: str, raw: Optional[str]) -> None:
|
||||
if raw is None:
|
||||
return
|
||||
values = [item.strip() for item in raw.split(",") if item.strip()]
|
||||
setattr(target, field_name, values)
|
||||
|
||||
|
||||
def _maybe_set_float(target: Any, field_name: str, raw: Optional[str]) -> None:
|
||||
if raw is None or not raw.strip():
|
||||
return
|
||||
try:
|
||||
setattr(target, field_name, float(raw))
|
||||
except ValueError:
|
||||
logger.warning("Invalid float for %s: %r", field_name, raw)
|
||||
|
||||
|
||||
def _maybe_set_str(target: Any, field_name: str, raw: Optional[str]) -> None:
|
||||
if raw is None:
|
||||
return
|
||||
value = raw.strip()
|
||||
if value:
|
||||
setattr(target, field_name, value)
|
||||
577
openspace/communication/gateway.py
Normal file
577
openspace/communication/gateway.py
Normal file
|
|
@ -0,0 +1,577 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import requests
|
||||
from aiohttp import web
|
||||
|
||||
from openspace.communication.adapters import FeishuAdapter, WhatsAppAdapter
|
||||
from openspace.communication.adapters.base import BaseChannelAdapter
|
||||
from openspace.communication.attachment_cache import AttachmentCache
|
||||
from openspace.communication.config import CommunicationConfig, load_communication_config
|
||||
from openspace.communication.gateway_runtime import RuntimeStatusStore, ScopedLock, ScopedLockManager
|
||||
from openspace.communication.policy import (
|
||||
build_attachment_instruction,
|
||||
is_authorized,
|
||||
should_accept_message,
|
||||
)
|
||||
from openspace.communication.runtime_manager import SessionRuntimeManager
|
||||
from openspace.communication.session_store import SessionStore
|
||||
from openspace.communication.types import ChannelMessage, ChannelPlatform, ChannelSession
|
||||
from openspace.host_detection import build_grounding_config_path, build_llm_kwargs, load_runtime_env
|
||||
from openspace.tool_layer import OpenSpace, OpenSpaceConfig
|
||||
from openspace.utils.logging import Logger
|
||||
|
||||
logger = Logger.get_logger(__name__)
|
||||
|
||||
|
||||
def _append_no_proxy_hosts(*hosts: str) -> None:
|
||||
for env_name in ("NO_PROXY", "no_proxy"):
|
||||
current = os.environ.get(env_name, "")
|
||||
entries = [entry.strip() for entry in current.split(",") if entry.strip()]
|
||||
updated = False
|
||||
for host in hosts:
|
||||
if host not in entries:
|
||||
entries.append(host)
|
||||
updated = True
|
||||
if updated:
|
||||
os.environ[env_name] = ",".join(entries)
|
||||
|
||||
|
||||
def _configure_ollama_process_env(model: str) -> None:
|
||||
if not model.lower().startswith("ollama/"):
|
||||
return
|
||||
|
||||
_append_no_proxy_hosts("127.0.0.1", "localhost")
|
||||
for env_name in (
|
||||
"HTTP_PROXY",
|
||||
"HTTPS_PROXY",
|
||||
"ALL_PROXY",
|
||||
"http_proxy",
|
||||
"https_proxy",
|
||||
"all_proxy",
|
||||
):
|
||||
if os.environ.get(env_name):
|
||||
logger.info("Clearing %s for local Ollama access", env_name)
|
||||
os.environ.pop(env_name, None)
|
||||
|
||||
|
||||
class CommunicationGateway:
|
||||
def __init__(self, config: CommunicationConfig):
|
||||
self.config = config
|
||||
workspace_root = (
|
||||
Path(config.agent.workspace_root).expanduser().resolve()
|
||||
if config.agent.workspace_root
|
||||
else None
|
||||
)
|
||||
self.session_store = SessionStore(
|
||||
config.sessions_dir,
|
||||
workspace_root=workspace_root,
|
||||
)
|
||||
self.attachment_cache = AttachmentCache(
|
||||
config.sessions_dir,
|
||||
max_attachment_bytes=config.sessions.max_attachment_bytes,
|
||||
max_session_attachment_bytes=config.sessions.max_session_attachment_bytes,
|
||||
)
|
||||
self.runtime_manager = SessionRuntimeManager(config, self._create_openspace_runtime)
|
||||
self._session_queues: Dict[str, asyncio.Queue[ChannelMessage]] = {}
|
||||
self._session_workers: Dict[str, asyncio.Task] = {}
|
||||
self._adapters: Dict[ChannelPlatform, BaseChannelAdapter] = {}
|
||||
self._web_app: Optional[web.Application] = None
|
||||
self._web_runner: Optional[web.AppRunner] = None
|
||||
self._web_site: Optional[web.TCPSite] = None
|
||||
self._running = False
|
||||
self._runtime_manager_started = False
|
||||
self._runtime_status = RuntimeStatusStore(self._runtime_status_path)
|
||||
self._lock_manager = ScopedLockManager(self._locks_dir)
|
||||
self._acquired_locks: list[ScopedLock] = []
|
||||
|
||||
async def start(self) -> None:
|
||||
if self._running:
|
||||
return
|
||||
|
||||
self.config.data_path.mkdir(parents=True, exist_ok=True)
|
||||
self._locks_dir.mkdir(parents=True, exist_ok=True)
|
||||
self._bridge_tokens_dir.mkdir(parents=True, exist_ok=True)
|
||||
self._outbound_media_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
try:
|
||||
self._build_adapters()
|
||||
for adapter in self._adapters.values():
|
||||
validate_configuration = getattr(adapter, "validate_configuration", None)
|
||||
if callable(validate_configuration):
|
||||
validate_configuration()
|
||||
self._acquire_adapter_locks()
|
||||
self._write_runtime_status("starting")
|
||||
await self.runtime_manager.start()
|
||||
self._runtime_manager_started = True
|
||||
|
||||
self._web_app = web.Application()
|
||||
self._web_app.router.add_get(self.config.server.health_path, self._handle_health)
|
||||
for adapter in self._adapters.values():
|
||||
adapter.register_http_routes(self._web_app)
|
||||
|
||||
self._web_runner = web.AppRunner(self._web_app)
|
||||
await self._web_runner.setup()
|
||||
self._web_site = web.TCPSite(
|
||||
self._web_runner,
|
||||
self.config.server.host,
|
||||
self.config.server.port,
|
||||
)
|
||||
await self._web_site.start()
|
||||
|
||||
for adapter in self._adapters.values():
|
||||
connected = await adapter.connect()
|
||||
if not connected:
|
||||
raise RuntimeError(
|
||||
f"Communication adapter failed to connect: {adapter.platform.value}"
|
||||
)
|
||||
|
||||
self._running = True
|
||||
self._write_runtime_status("running")
|
||||
logger.info(
|
||||
"Communication gateway started on %s:%s for platforms=%s",
|
||||
self.config.server.host,
|
||||
self.config.server.port,
|
||||
",".join(self.config.enabled_platforms) or "(none)",
|
||||
)
|
||||
except Exception as exc:
|
||||
await self._rollback_start(exc)
|
||||
raise
|
||||
|
||||
async def stop(self) -> None:
|
||||
if not self._running and not self._has_live_resources():
|
||||
return
|
||||
|
||||
self._write_runtime_status("stopping")
|
||||
self._running = False
|
||||
await self._stop_session_workers()
|
||||
await self._disconnect_adapters()
|
||||
await self._cleanup_web_runner()
|
||||
await self._stop_runtime_manager()
|
||||
self._release_locks()
|
||||
self._write_runtime_status("stopped")
|
||||
logger.info("Communication gateway stopped")
|
||||
|
||||
def _build_adapters(self) -> None:
|
||||
adapters: Dict[ChannelPlatform, BaseChannelAdapter] = {}
|
||||
if self.config.whatsapp.enabled:
|
||||
adapter = self._instantiate_adapter(
|
||||
WhatsAppAdapter,
|
||||
self.config.whatsapp,
|
||||
self.attachment_cache,
|
||||
runtime_dir=self.config.data_path,
|
||||
poll_interval_seconds=self.config.sessions.whatsapp_poll_interval_seconds,
|
||||
)
|
||||
adapter.set_message_handler(self.handle_message)
|
||||
adapters[ChannelPlatform.WHATSAPP] = adapter
|
||||
if self.config.feishu.enabled:
|
||||
adapter = self._instantiate_adapter(
|
||||
FeishuAdapter,
|
||||
self.config.feishu,
|
||||
self.attachment_cache,
|
||||
runtime_dir=self.config.data_path,
|
||||
)
|
||||
adapter.set_message_handler(self.handle_message)
|
||||
adapters[ChannelPlatform.FEISHU] = adapter
|
||||
self._adapters = adapters
|
||||
|
||||
@staticmethod
|
||||
def _instantiate_adapter(adapter_cls: Any, *args: Any, **kwargs: Any) -> BaseChannelAdapter:
|
||||
try:
|
||||
return adapter_cls(*args, **kwargs)
|
||||
except TypeError as exc:
|
||||
if "unexpected keyword argument" not in str(exc):
|
||||
raise
|
||||
compatibility_kwargs = dict(kwargs)
|
||||
compatibility_kwargs.pop("runtime_dir", None)
|
||||
return adapter_cls(*args, **compatibility_kwargs)
|
||||
|
||||
def _acquire_adapter_locks(self) -> None:
|
||||
self._release_locks()
|
||||
for adapter in self._adapters.values():
|
||||
get_lock_identity = getattr(adapter, "get_lock_identity", None)
|
||||
binding = get_lock_identity() if callable(get_lock_identity) else None
|
||||
if binding is None:
|
||||
continue
|
||||
scope, identity = binding
|
||||
lock = self._lock_manager.acquire(
|
||||
scope=scope,
|
||||
identity=identity,
|
||||
metadata={"platform": adapter.platform.value},
|
||||
)
|
||||
self._acquired_locks.append(lock)
|
||||
|
||||
def _release_locks(self) -> None:
|
||||
while self._acquired_locks:
|
||||
self._lock_manager.release(self._acquired_locks.pop())
|
||||
|
||||
def _write_runtime_status(
|
||||
self,
|
||||
gateway_state: str,
|
||||
*,
|
||||
fatal_error: Optional[str] = None,
|
||||
) -> None:
|
||||
platform_states = {
|
||||
adapter.platform.value: {"connected": adapter.is_connected}
|
||||
for adapter in self._adapters.values()
|
||||
}
|
||||
self._runtime_status.write(
|
||||
gateway_state=gateway_state,
|
||||
platforms=platform_states,
|
||||
config_path=str(self.config.data_path),
|
||||
fatal_error=fatal_error,
|
||||
)
|
||||
|
||||
async def _rollback_start(self, exc: Exception) -> None:
|
||||
logger.error("Communication gateway startup failed: %s", exc, exc_info=True)
|
||||
self._running = False
|
||||
await self._disconnect_adapters()
|
||||
await self._cleanup_web_runner()
|
||||
try:
|
||||
await self._stop_runtime_manager()
|
||||
finally:
|
||||
self._release_locks()
|
||||
self._write_runtime_status("failed", fatal_error=str(exc))
|
||||
|
||||
async def _stop_session_workers(self) -> None:
|
||||
worker_tasks = list(self._session_workers.values())
|
||||
self._session_workers.clear()
|
||||
for task in worker_tasks:
|
||||
task.cancel()
|
||||
if worker_tasks:
|
||||
await asyncio.gather(*worker_tasks, return_exceptions=True)
|
||||
self._session_queues.clear()
|
||||
|
||||
async def _disconnect_adapters(self) -> None:
|
||||
adapters = list(self._adapters.values())
|
||||
self._adapters.clear()
|
||||
for adapter in adapters:
|
||||
try:
|
||||
await adapter.disconnect()
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to disconnect adapter during cleanup: %s",
|
||||
getattr(adapter.platform, "value", "unknown"),
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
async def _cleanup_web_runner(self) -> None:
|
||||
if self._web_runner is None:
|
||||
return
|
||||
try:
|
||||
await self._web_runner.cleanup()
|
||||
finally:
|
||||
self._web_runner = None
|
||||
self._web_site = None
|
||||
self._web_app = None
|
||||
|
||||
async def _stop_runtime_manager(self) -> None:
|
||||
if not self._runtime_manager_started:
|
||||
return
|
||||
try:
|
||||
await self.runtime_manager.stop()
|
||||
finally:
|
||||
self._runtime_manager_started = False
|
||||
|
||||
def _has_live_resources(self) -> bool:
|
||||
return any(
|
||||
(
|
||||
self._runtime_manager_started,
|
||||
bool(self._adapters),
|
||||
self._web_runner is not None,
|
||||
bool(self._acquired_locks),
|
||||
bool(self._session_workers),
|
||||
bool(self._session_queues),
|
||||
)
|
||||
)
|
||||
|
||||
@property
|
||||
def _runtime_status_path(self) -> Path:
|
||||
return getattr(self.config, "runtime_status_path", self.config.data_path / "runtime_status.json")
|
||||
|
||||
@property
|
||||
def _locks_dir(self) -> Path:
|
||||
return getattr(self.config, "locks_dir", self.config.data_path / "locks")
|
||||
|
||||
@property
|
||||
def _bridge_tokens_dir(self) -> Path:
|
||||
return getattr(self.config, "bridge_tokens_dir", self.config.data_path / "bridge_tokens")
|
||||
|
||||
@property
|
||||
def _outbound_media_dir(self) -> Path:
|
||||
return getattr(self.config, "outbound_media_dir", self.config.data_path / "outbound_media")
|
||||
|
||||
async def handle_message(self, message: ChannelMessage) -> None:
|
||||
session = self.session_store.get_or_create_session(message.source)
|
||||
queue = self._session_queues.get(session.session_key)
|
||||
if queue is None:
|
||||
queue = asyncio.Queue(maxsize=self.config.sessions.per_session_queue_size)
|
||||
self._session_queues[session.session_key] = queue
|
||||
worker = self._session_workers.get(session.session_key)
|
||||
if worker is None or worker.done():
|
||||
self._session_workers[session.session_key] = asyncio.create_task(
|
||||
self._session_worker(session, queue)
|
||||
)
|
||||
await queue.put(message)
|
||||
|
||||
async def _session_worker(
|
||||
self,
|
||||
session: ChannelSession,
|
||||
queue: asyncio.Queue[ChannelMessage],
|
||||
) -> None:
|
||||
session_key = session.session_key
|
||||
try:
|
||||
while True:
|
||||
try:
|
||||
message = await asyncio.wait_for(
|
||||
queue.get(),
|
||||
timeout=self.config.sessions.idle_ttl_seconds,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
if queue.empty():
|
||||
logger.info("Retiring idle communication worker: %s", session_key)
|
||||
return
|
||||
continue
|
||||
|
||||
try:
|
||||
await self._process_message(session, message)
|
||||
except Exception as exc:
|
||||
logger.error(
|
||||
"Failed to process %s message for session %s: %s",
|
||||
message.source.platform.value,
|
||||
session.session_key,
|
||||
exc,
|
||||
exc_info=True,
|
||||
)
|
||||
adapter = self._adapters.get(message.source.platform)
|
||||
if adapter:
|
||||
await adapter.send_text(
|
||||
message.source.chat_id,
|
||||
f"OpenSpace communication error: {exc}",
|
||||
)
|
||||
finally:
|
||||
queue.task_done()
|
||||
finally:
|
||||
current_task = asyncio.current_task()
|
||||
if self._session_workers.get(session_key) is current_task:
|
||||
self._session_workers.pop(session_key, None)
|
||||
if queue.empty():
|
||||
if (
|
||||
self._session_queues.get(session_key) is queue
|
||||
and session_key not in self._session_workers
|
||||
):
|
||||
self._session_queues.pop(session_key, None)
|
||||
elif self._running and session_key not in self._session_workers:
|
||||
self._session_workers[session_key] = asyncio.create_task(
|
||||
self._session_worker(session, queue)
|
||||
)
|
||||
|
||||
async def _process_message(self, session: ChannelSession, message: ChannelMessage) -> None:
|
||||
platform_config = self._get_platform_config(message.source.platform)
|
||||
if not is_authorized(message, platform_config):
|
||||
logger.info(
|
||||
"Rejected %s message from unauthorized user %s",
|
||||
message.source.platform.value,
|
||||
message.source.user_id,
|
||||
)
|
||||
return
|
||||
|
||||
reply_to_bot = self.session_store.is_reply_to_assistant(
|
||||
session,
|
||||
message.reply_to_message_id,
|
||||
)
|
||||
if not should_accept_message(message, platform_config, reply_to_bot):
|
||||
logger.debug(
|
||||
"Skipped %s group message that did not satisfy policy",
|
||||
message.source.platform.value,
|
||||
)
|
||||
return
|
||||
|
||||
history = self.session_store.load_history(
|
||||
session,
|
||||
self.config.sessions.history_max_turns,
|
||||
)
|
||||
if not message.text.strip():
|
||||
message.text = build_attachment_instruction(message)
|
||||
|
||||
self.session_store.append_user_message(session, message)
|
||||
|
||||
result = await self.runtime_manager.execute_turn(
|
||||
session=session,
|
||||
message=message,
|
||||
conversation_history=history,
|
||||
channel_context=message.to_channel_context(session.session_key),
|
||||
)
|
||||
response_text = self._extract_response_text(result)
|
||||
|
||||
adapter = self._adapters.get(message.source.platform)
|
||||
if adapter is None:
|
||||
raise RuntimeError(f"No adapter registered for {message.source.platform.value}")
|
||||
|
||||
send_result = await adapter.send_text(
|
||||
message.source.chat_id,
|
||||
response_text,
|
||||
reply_to_message_id=message.message_id,
|
||||
)
|
||||
if not send_result.success:
|
||||
logger.warning(
|
||||
"Failed to send %s response for session %s: %s",
|
||||
message.source.platform.value,
|
||||
session.session_key,
|
||||
send_result.error,
|
||||
)
|
||||
self.session_store.append_assistant_message(
|
||||
session,
|
||||
content=response_text,
|
||||
platform_message_id=send_result.message_id,
|
||||
metadata={
|
||||
"task_id": result.get("task_id"),
|
||||
"status": result.get("status"),
|
||||
"send_success": send_result.success,
|
||||
"send_error": send_result.error,
|
||||
},
|
||||
)
|
||||
|
||||
async def _handle_health(self, request: web.Request) -> web.Response:
|
||||
runtime_status = await self.runtime_manager.status()
|
||||
gateway_status = self._runtime_status.read() or {}
|
||||
return web.json_response(
|
||||
{
|
||||
"status": "ok" if self._running else "starting",
|
||||
"gateway": gateway_status,
|
||||
"platforms": {
|
||||
platform.value: {
|
||||
"connected": adapter.is_connected,
|
||||
}
|
||||
for platform, adapter in self._adapters.items()
|
||||
},
|
||||
"runtime": runtime_status,
|
||||
"sessions": len(self.session_store.list_sessions()),
|
||||
}
|
||||
)
|
||||
|
||||
async def _create_openspace_runtime(self, session: ChannelSession) -> OpenSpace:
|
||||
load_runtime_env()
|
||||
env_model = os.environ.get("OPENSPACE_MODEL", "")
|
||||
model, llm_kwargs = build_llm_kwargs(env_model)
|
||||
llm_kwargs = dict(llm_kwargs)
|
||||
if model.lower().startswith("ollama/"):
|
||||
llm_kwargs["api_base"] = os.environ.get("OLLAMA_API_BASE", "").strip() or "http://127.0.0.1:11434"
|
||||
llm_kwargs["api_key"] = os.environ.get("OLLAMA_API_KEY", "").strip() or llm_kwargs.get("api_key") or "ollama"
|
||||
llm_kwargs.pop("extra_headers", None)
|
||||
backend_scope = self.config.agent.backend_scope
|
||||
grounding_config_path = (
|
||||
self.config.agent.grounding_config_path
|
||||
or build_grounding_config_path()
|
||||
)
|
||||
recording_dir = self.config.data_path / "recordings"
|
||||
openspace_config = OpenSpaceConfig(
|
||||
llm_model=model,
|
||||
llm_kwargs=llm_kwargs,
|
||||
workspace_dir=session.workspace_dir,
|
||||
grounding_max_iterations=self.config.agent.max_iterations,
|
||||
enable_recording=self.config.agent.enable_recording,
|
||||
recording_backends=self.config.agent.recording_backends,
|
||||
recording_log_dir=str(recording_dir),
|
||||
backend_scope=backend_scope,
|
||||
grounding_config_path=grounding_config_path,
|
||||
llm_timeout=self.config.agent.llm_timeout,
|
||||
)
|
||||
runtime = OpenSpace(openspace_config)
|
||||
await runtime.initialize()
|
||||
return runtime
|
||||
|
||||
def _get_platform_config(self, platform: ChannelPlatform) -> Any:
|
||||
if platform == ChannelPlatform.WHATSAPP:
|
||||
return self.config.whatsapp
|
||||
if platform == ChannelPlatform.FEISHU:
|
||||
return self.config.feishu
|
||||
raise ValueError(f"Unsupported platform: {platform}")
|
||||
|
||||
@staticmethod
|
||||
def _extract_response_text(result: Dict[str, Any]) -> str:
|
||||
response = str(result.get("response", "")).strip()
|
||||
if response:
|
||||
return response
|
||||
error = str(result.get("error", "")).strip()
|
||||
if error:
|
||||
return f"OpenSpace error: {error}"
|
||||
return "OpenSpace completed the task but returned no response."
|
||||
|
||||
|
||||
def _build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="OpenSpace communication gateway",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--config",
|
||||
type=str,
|
||||
help="Path to the communication JSON config file",
|
||||
)
|
||||
subparsers = parser.add_subparsers(dest="command")
|
||||
run_parser = subparsers.add_parser("run", help="Start the communication gateway")
|
||||
run_parser.add_argument(
|
||||
"--config",
|
||||
type=str,
|
||||
help="Path to the communication JSON config file",
|
||||
)
|
||||
health_parser = subparsers.add_parser("health", help="Check the running gateway health endpoint")
|
||||
health_parser.add_argument(
|
||||
"--config",
|
||||
type=str,
|
||||
help="Path to the communication JSON config file",
|
||||
)
|
||||
health_parser.add_argument("--host", type=str, default=None)
|
||||
health_parser.add_argument("--port", type=int, default=None)
|
||||
return parser
|
||||
|
||||
|
||||
async def _run_gateway(config_path: Optional[str]) -> int:
|
||||
config = load_communication_config(config_path)
|
||||
_configure_ollama_process_env(os.environ.get("OPENSPACE_MODEL", ""))
|
||||
gateway = CommunicationGateway(config)
|
||||
try:
|
||||
await gateway.start()
|
||||
except Exception as exc:
|
||||
logger.error("Failed to start communication gateway: %s", exc)
|
||||
return 1
|
||||
|
||||
try:
|
||||
while True:
|
||||
await asyncio.sleep(3600)
|
||||
except (asyncio.CancelledError, KeyboardInterrupt):
|
||||
pass
|
||||
finally:
|
||||
await gateway.stop()
|
||||
return 0
|
||||
|
||||
|
||||
def _check_health(config_path: Optional[str], host: Optional[str], port: Optional[int]) -> int:
|
||||
config = load_communication_config(config_path)
|
||||
url = f"http://{host or config.server.host}:{port or config.server.port}{config.server.health_path}"
|
||||
response = requests.get(url, timeout=5)
|
||||
response.raise_for_status()
|
||||
print(response.text)
|
||||
return 0
|
||||
|
||||
|
||||
async def main(argv: Optional[list[str]] = None) -> int:
|
||||
parser = _build_parser()
|
||||
args = parser.parse_args(argv)
|
||||
command = args.command or "run"
|
||||
if command == "health":
|
||||
return _check_health(args.config, args.host, args.port)
|
||||
return await _run_gateway(args.config)
|
||||
|
||||
|
||||
def run_main() -> None:
|
||||
raise SystemExit(asyncio.run(main()))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
run_main()
|
||||
252
openspace/communication/gateway_runtime.py
Normal file
252
openspace/communication/gateway_runtime.py
Normal file
|
|
@ -0,0 +1,252 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
from openspace.utils.logging import Logger
|
||||
|
||||
logger = Logger.get_logger(__name__)
|
||||
|
||||
_GATEWAY_KIND = "openspace-communication-gateway"
|
||||
|
||||
|
||||
def _utcnow_iso() -> str:
|
||||
return datetime.now(timezone.utc).isoformat()
|
||||
|
||||
|
||||
def _scope_hash(identity: str) -> str:
|
||||
return hashlib.sha256(identity.encode("utf-8")).hexdigest()[:16]
|
||||
|
||||
|
||||
def _process_start_time(pid: int) -> Optional[int]:
|
||||
stat_path = Path(f"/proc/{pid}/stat")
|
||||
try:
|
||||
return int(stat_path.read_text(encoding="utf-8").split()[21])
|
||||
except (FileNotFoundError, IndexError, PermissionError, ValueError, OSError):
|
||||
return None
|
||||
|
||||
|
||||
def _is_pid_alive(pid: int) -> bool:
|
||||
if pid <= 0:
|
||||
return False
|
||||
try:
|
||||
os.kill(pid, 0)
|
||||
except ProcessLookupError:
|
||||
return False
|
||||
except PermissionError:
|
||||
return True
|
||||
else:
|
||||
return True
|
||||
|
||||
|
||||
def _build_process_record() -> dict[str, Any]:
|
||||
pid = os.getpid()
|
||||
return {
|
||||
"pid": pid,
|
||||
"kind": _GATEWAY_KIND,
|
||||
"argv": list(sys.argv),
|
||||
"start_time": _process_start_time(pid),
|
||||
}
|
||||
|
||||
|
||||
def _record_matches_live_process(record: dict[str, Any]) -> bool:
|
||||
try:
|
||||
pid = int(record["pid"])
|
||||
except (KeyError, TypeError, ValueError):
|
||||
return False
|
||||
if not _is_pid_alive(pid):
|
||||
return False
|
||||
|
||||
recorded_start_time = record.get("start_time")
|
||||
live_start_time = _process_start_time(pid)
|
||||
if recorded_start_time is None or live_start_time is None:
|
||||
return True
|
||||
return live_start_time == recorded_start_time
|
||||
|
||||
|
||||
def _read_json(path: Path) -> Optional[dict[str, Any]]:
|
||||
if not path.exists():
|
||||
return None
|
||||
try:
|
||||
payload = json.loads(path.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError):
|
||||
return None
|
||||
return payload if isinstance(payload, dict) else None
|
||||
|
||||
|
||||
def _write_json(path: Path, payload: dict[str, Any]) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
|
||||
|
||||
class LockConflictError(RuntimeError):
|
||||
def __init__(self, scope: str, identity: str, record: Optional[dict[str, Any]] = None):
|
||||
super().__init__(f"Communication gateway lock is already held for {scope}:{identity}")
|
||||
self.scope = scope
|
||||
self.identity = identity
|
||||
self.record = record or {}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ScopedRuntimeLock:
|
||||
path: Path
|
||||
record: dict[str, Any]
|
||||
released: bool = False
|
||||
|
||||
@classmethod
|
||||
def acquire(
|
||||
cls,
|
||||
*,
|
||||
locks_dir: Path,
|
||||
scope: str,
|
||||
identity: str,
|
||||
metadata: Optional[dict[str, Any]] = None,
|
||||
) -> "ScopedRuntimeLock":
|
||||
locks_dir.mkdir(parents=True, exist_ok=True)
|
||||
lock_path = locks_dir / f"{scope}-{_scope_hash(identity)}.lock"
|
||||
record = {
|
||||
**_build_process_record(),
|
||||
"scope": scope,
|
||||
"identity": identity,
|
||||
"metadata": metadata or {},
|
||||
"created_at": _utcnow_iso(),
|
||||
"updated_at": _utcnow_iso(),
|
||||
}
|
||||
|
||||
while True:
|
||||
existing = _read_json(lock_path)
|
||||
if existing is not None:
|
||||
if _record_matches_live_process(existing):
|
||||
raise LockConflictError(scope, identity, existing)
|
||||
try:
|
||||
lock_path.unlink()
|
||||
except FileNotFoundError:
|
||||
continue
|
||||
except OSError as exc:
|
||||
raise RuntimeError(f"Failed to remove stale gateway lock {lock_path}: {exc}") from exc
|
||||
elif lock_path.exists():
|
||||
try:
|
||||
lock_path.unlink()
|
||||
except FileNotFoundError:
|
||||
continue
|
||||
except OSError as exc:
|
||||
raise RuntimeError(
|
||||
f"Failed to remove malformed gateway lock {lock_path}: {exc}"
|
||||
) from exc
|
||||
continue
|
||||
|
||||
try:
|
||||
fd = os.open(lock_path, os.O_CREAT | os.O_EXCL | os.O_WRONLY, 0o600)
|
||||
except FileExistsError:
|
||||
continue
|
||||
|
||||
try:
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as handle:
|
||||
json.dump(record, handle, ensure_ascii=False, indent=2)
|
||||
except Exception:
|
||||
try:
|
||||
lock_path.unlink(missing_ok=True)
|
||||
except OSError:
|
||||
pass
|
||||
raise
|
||||
return cls(path=lock_path, record=record)
|
||||
|
||||
def release(self) -> None:
|
||||
if self.released:
|
||||
return
|
||||
try:
|
||||
current = _read_json(self.path)
|
||||
if current and current.get("pid") == self.record.get("pid"):
|
||||
self.path.unlink(missing_ok=True)
|
||||
except OSError as exc:
|
||||
logger.warning("Failed to release communication gateway lock %s: %s", self.path, exc)
|
||||
finally:
|
||||
self.released = True
|
||||
|
||||
|
||||
class GatewayRuntimeTracker:
|
||||
def __init__(self, status_path: Path):
|
||||
self.status_path = status_path
|
||||
|
||||
def write_status(
|
||||
self,
|
||||
*,
|
||||
gateway_state: str,
|
||||
platforms: dict[str, dict[str, Any]],
|
||||
fatal_error: Optional[str] = None,
|
||||
exit_reason: Optional[str] = None,
|
||||
config_path: Optional[str] = None,
|
||||
sessions: Optional[int] = None,
|
||||
) -> None:
|
||||
payload = {
|
||||
**_build_process_record(),
|
||||
"gateway_state": gateway_state,
|
||||
"fatal_error": fatal_error,
|
||||
"exit_reason": exit_reason,
|
||||
"config_path": config_path,
|
||||
"platforms": platforms,
|
||||
"sessions": sessions,
|
||||
"updated_at": _utcnow_iso(),
|
||||
}
|
||||
_write_json(self.status_path, payload)
|
||||
|
||||
def read_status(self) -> Optional[dict[str, Any]]:
|
||||
return _read_json(self.status_path)
|
||||
|
||||
|
||||
class ScopedLockManager:
|
||||
def __init__(self, locks_dir: Path):
|
||||
self.locks_dir = locks_dir
|
||||
|
||||
def acquire(
|
||||
self,
|
||||
scope: str,
|
||||
identity: str,
|
||||
metadata: Optional[dict[str, Any]] = None,
|
||||
) -> "ScopedLock":
|
||||
return ScopedRuntimeLock.acquire(
|
||||
locks_dir=self.locks_dir,
|
||||
scope=scope,
|
||||
identity=identity,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def release(lock: "ScopedLock") -> None:
|
||||
lock.release()
|
||||
|
||||
|
||||
class RuntimeStatusStore:
|
||||
def __init__(self, status_path: Path):
|
||||
self._tracker = GatewayRuntimeTracker(status_path)
|
||||
|
||||
def write(
|
||||
self,
|
||||
*,
|
||||
gateway_state: str,
|
||||
platforms: dict[str, dict[str, Any]],
|
||||
fatal_error: Optional[str] = None,
|
||||
exit_reason: Optional[str] = None,
|
||||
config_path: Optional[str] = None,
|
||||
sessions: Optional[int] = None,
|
||||
) -> None:
|
||||
self._tracker.write_status(
|
||||
gateway_state=gateway_state,
|
||||
platforms=platforms,
|
||||
fatal_error=fatal_error,
|
||||
exit_reason=exit_reason,
|
||||
config_path=config_path,
|
||||
sessions=sessions,
|
||||
)
|
||||
|
||||
def read(self) -> Optional[dict[str, Any]]:
|
||||
return self._tracker.read_status()
|
||||
|
||||
|
||||
ScopedLock = ScopedRuntimeLock
|
||||
61
openspace/communication/policy.py
Normal file
61
openspace/communication/policy.py
Normal file
|
|
@ -0,0 +1,61 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from openspace.communication.types import ChannelMessage
|
||||
|
||||
|
||||
def build_attachment_instruction(message: ChannelMessage) -> str:
|
||||
attachment_paths = ", ".join(attachment.path for attachment in message.attachments)
|
||||
return (
|
||||
"Please inspect the attached files and help the user based on their contents. "
|
||||
f"Attachment paths: {attachment_paths}"
|
||||
)
|
||||
|
||||
|
||||
def is_authorized(message: ChannelMessage, platform_config: Any) -> bool:
|
||||
if getattr(platform_config, "allow_all_users", False):
|
||||
return True
|
||||
|
||||
allowed_users = {
|
||||
entry.strip()
|
||||
for entry in getattr(platform_config, "allowed_users", [])
|
||||
if entry and entry.strip()
|
||||
}
|
||||
if not allowed_users:
|
||||
return False
|
||||
|
||||
user_candidates = {
|
||||
candidate.strip()
|
||||
for candidate in (
|
||||
message.source.user_id,
|
||||
message.source.user_name,
|
||||
message.metadata.get("raw_user_id") if isinstance(message.metadata, dict) else None,
|
||||
)
|
||||
if candidate and candidate.strip()
|
||||
}
|
||||
if isinstance(message.metadata, dict):
|
||||
for candidate in message.metadata.get("auth_candidates", []) or []:
|
||||
if isinstance(candidate, str) and candidate.strip():
|
||||
user_candidates.add(candidate.strip())
|
||||
return bool(user_candidates & allowed_users)
|
||||
|
||||
|
||||
def should_accept_message(
|
||||
message: ChannelMessage,
|
||||
platform_config: Any,
|
||||
reply_to_bot: bool,
|
||||
) -> bool:
|
||||
if message.source.chat_type == "dm":
|
||||
return bool(getattr(platform_config, "allow_dm", True))
|
||||
if not getattr(platform_config, "allow_groups", True):
|
||||
return False
|
||||
|
||||
group_policy = str(getattr(platform_config, "group_policy", "reply_or_mention"))
|
||||
if group_policy == "disabled":
|
||||
return False
|
||||
if group_policy == "all":
|
||||
return True
|
||||
if group_policy == "mention_only":
|
||||
return message.mentions_bot
|
||||
return message.mentions_bot or reply_to_bot
|
||||
149
openspace/communication/runtime_manager.py
Normal file
149
openspace/communication/runtime_manager.py
Normal file
|
|
@ -0,0 +1,149 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import uuid
|
||||
from time import monotonic
|
||||
from typing import Any, Awaitable, Callable, Dict, Optional
|
||||
|
||||
from openspace.tool_layer import OpenSpace
|
||||
from openspace.utils.logging import Logger
|
||||
|
||||
from .config import CommunicationConfig
|
||||
from .types import ChannelMessage, ChannelSession
|
||||
|
||||
logger = Logger.get_logger(__name__)
|
||||
|
||||
|
||||
OpenSpaceFactory = Callable[[ChannelSession], Awaitable[OpenSpace]]
|
||||
|
||||
|
||||
class SessionRuntime:
|
||||
def __init__(self, session: ChannelSession, openspace_factory: OpenSpaceFactory):
|
||||
self.session = session
|
||||
self._openspace_factory = openspace_factory
|
||||
self._openspace: Optional[OpenSpace] = None
|
||||
self._lock = asyncio.Lock()
|
||||
self.last_used_monotonic = monotonic()
|
||||
|
||||
@property
|
||||
def openspace(self) -> Optional[OpenSpace]:
|
||||
return self._openspace
|
||||
|
||||
async def ensure_initialized(self) -> OpenSpace:
|
||||
if self._openspace is None:
|
||||
self._openspace = await self._openspace_factory(self.session)
|
||||
self.last_used_monotonic = monotonic()
|
||||
return self._openspace
|
||||
|
||||
async def execute_turn(
|
||||
self,
|
||||
*,
|
||||
message: ChannelMessage,
|
||||
conversation_history: list[dict[str, str]],
|
||||
channel_context: dict[str, Any],
|
||||
max_iterations: Optional[int] = None,
|
||||
) -> Dict[str, Any]:
|
||||
async with self._lock:
|
||||
openspace = await self.ensure_initialized()
|
||||
self.last_used_monotonic = monotonic()
|
||||
task_id = f"comm_{self.session.session_key}_{uuid.uuid4().hex[:10]}"
|
||||
result = await openspace.execute(
|
||||
task=message.text,
|
||||
context={
|
||||
"conversation_history": conversation_history,
|
||||
"channel_context": channel_context,
|
||||
"session_key": self.session.session_key,
|
||||
},
|
||||
workspace_dir=self.session.workspace_dir,
|
||||
max_iterations=max_iterations,
|
||||
task_id=task_id,
|
||||
)
|
||||
self.last_used_monotonic = monotonic()
|
||||
return result
|
||||
|
||||
async def close(self) -> None:
|
||||
if self._openspace is not None:
|
||||
await self._openspace.cleanup()
|
||||
self._openspace = None
|
||||
|
||||
def is_idle(self, idle_ttl_seconds: int) -> bool:
|
||||
if self._lock.locked():
|
||||
return False
|
||||
return (monotonic() - self.last_used_monotonic) >= idle_ttl_seconds
|
||||
|
||||
|
||||
class SessionRuntimeManager:
|
||||
def __init__(
|
||||
self,
|
||||
config: CommunicationConfig,
|
||||
openspace_factory: OpenSpaceFactory,
|
||||
):
|
||||
self.config = config
|
||||
self._openspace_factory = openspace_factory
|
||||
self._runtimes: Dict[str, SessionRuntime] = {}
|
||||
self._lock = asyncio.Lock()
|
||||
self._semaphore = asyncio.Semaphore(config.sessions.max_parallel_sessions)
|
||||
self._eviction_task: Optional[asyncio.Task] = None
|
||||
|
||||
async def start(self) -> None:
|
||||
if self._eviction_task is None:
|
||||
self._eviction_task = asyncio.create_task(self._evict_idle_loop())
|
||||
|
||||
async def stop(self) -> None:
|
||||
if self._eviction_task is not None:
|
||||
self._eviction_task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError):
|
||||
await self._eviction_task
|
||||
self._eviction_task = None
|
||||
|
||||
async with self._lock:
|
||||
runtimes = list(self._runtimes.values())
|
||||
self._runtimes.clear()
|
||||
for runtime in runtimes:
|
||||
await runtime.close()
|
||||
|
||||
async def execute_turn(
|
||||
self,
|
||||
*,
|
||||
session: ChannelSession,
|
||||
message: ChannelMessage,
|
||||
conversation_history: list[dict[str, str]],
|
||||
channel_context: dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
runtime = await self._get_or_create_runtime(session)
|
||||
async with self._semaphore:
|
||||
return await runtime.execute_turn(
|
||||
message=message,
|
||||
conversation_history=conversation_history,
|
||||
channel_context=channel_context,
|
||||
max_iterations=self.config.agent.max_iterations,
|
||||
)
|
||||
|
||||
async def status(self) -> Dict[str, Any]:
|
||||
async with self._lock:
|
||||
return {
|
||||
"active_runtimes": len(self._runtimes),
|
||||
"session_keys": sorted(self._runtimes.keys()),
|
||||
}
|
||||
|
||||
async def _get_or_create_runtime(self, session: ChannelSession) -> SessionRuntime:
|
||||
async with self._lock:
|
||||
runtime = self._runtimes.get(session.session_key)
|
||||
if runtime is None:
|
||||
runtime = SessionRuntime(session, self._openspace_factory)
|
||||
self._runtimes[session.session_key] = runtime
|
||||
return runtime
|
||||
|
||||
async def _evict_idle_loop(self) -> None:
|
||||
while True:
|
||||
await asyncio.sleep(30)
|
||||
stale_keys: list[str] = []
|
||||
async with self._lock:
|
||||
for session_key, runtime in self._runtimes.items():
|
||||
if runtime.is_idle(self.config.sessions.idle_ttl_seconds):
|
||||
stale_keys.append(session_key)
|
||||
runtimes = [self._runtimes.pop(key) for key in stale_keys]
|
||||
for runtime in runtimes:
|
||||
logger.info("Evicting idle communication runtime: %s", runtime.session.session_key)
|
||||
await runtime.close()
|
||||
199
openspace/communication/session_store.py
Normal file
199
openspace/communication/session_store.py
Normal file
|
|
@ -0,0 +1,199 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from openspace.utils.logging import Logger
|
||||
|
||||
from .types import ChannelAttachment, ChannelMessage, ChannelSession, ChannelSource
|
||||
|
||||
logger = Logger.get_logger(__name__)
|
||||
|
||||
|
||||
class SessionStore:
|
||||
def __init__(
|
||||
self,
|
||||
sessions_dir: Path,
|
||||
*,
|
||||
workspace_root: Optional[Path] = None,
|
||||
):
|
||||
self.sessions_dir = sessions_dir
|
||||
self.workspace_root = workspace_root
|
||||
self.sessions_dir.mkdir(parents=True, exist_ok=True)
|
||||
if self.workspace_root is not None:
|
||||
self.workspace_root.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def get_or_create_session(self, source: ChannelSource) -> ChannelSession:
|
||||
session_key = build_session_key(source)
|
||||
session_dir = self.sessions_dir / session_key
|
||||
session_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
metadata_path = session_dir / "session.json"
|
||||
transcript_path = session_dir / "transcript.jsonl"
|
||||
attachments_dir = session_dir / "attachments"
|
||||
workspace_dir = (
|
||||
self.workspace_root / session_key
|
||||
if self.workspace_root is not None
|
||||
else session_dir / "workspace"
|
||||
)
|
||||
attachments_dir.mkdir(parents=True, exist_ok=True)
|
||||
workspace_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
now = _utcnow_iso()
|
||||
if metadata_path.exists():
|
||||
with open(metadata_path, "r", encoding="utf-8") as handle:
|
||||
data = json.load(handle)
|
||||
session = ChannelSession.from_dict(data)
|
||||
session.source = source
|
||||
session.updated_at = now
|
||||
else:
|
||||
session = ChannelSession(
|
||||
session_key=session_key,
|
||||
source=source,
|
||||
session_dir=str(session_dir),
|
||||
workspace_dir=str(workspace_dir),
|
||||
attachments_dir=str(attachments_dir),
|
||||
transcript_path=str(transcript_path),
|
||||
metadata_path=str(metadata_path),
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
|
||||
self._write_session_metadata(session)
|
||||
return session
|
||||
|
||||
def append_user_message(self, session: ChannelSession, message: ChannelMessage) -> None:
|
||||
self._append_transcript_entry(
|
||||
session,
|
||||
{
|
||||
"entry_id": uuid.uuid4().hex,
|
||||
"role": "user",
|
||||
"content": message.text,
|
||||
"platform_message_id": message.message_id,
|
||||
"reply_to_message_id": message.reply_to_message_id,
|
||||
"reply_to_text": message.reply_to_text,
|
||||
"mentions_bot": message.mentions_bot,
|
||||
"attachments": [attachment.to_context_dict() for attachment in message.attachments],
|
||||
"source": message.source.to_dict(),
|
||||
"metadata": message.metadata,
|
||||
"timestamp": message.received_at.isoformat(),
|
||||
},
|
||||
)
|
||||
|
||||
def append_assistant_message(
|
||||
self,
|
||||
session: ChannelSession,
|
||||
*,
|
||||
content: str,
|
||||
platform_message_id: Optional[str] = None,
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
) -> None:
|
||||
self._append_transcript_entry(
|
||||
session,
|
||||
{
|
||||
"entry_id": uuid.uuid4().hex,
|
||||
"role": "assistant",
|
||||
"content": content,
|
||||
"platform_message_id": platform_message_id,
|
||||
"metadata": metadata or {},
|
||||
"timestamp": _utcnow_iso(),
|
||||
},
|
||||
)
|
||||
|
||||
def load_history(self, session: ChannelSession, max_turns: int) -> List[Dict[str, str]]:
|
||||
entries = self._read_transcript_entries(session)
|
||||
if not entries:
|
||||
return []
|
||||
|
||||
selected: List[Dict[str, str]] = []
|
||||
user_messages = 0
|
||||
for entry in reversed(entries):
|
||||
role = entry.get("role")
|
||||
if role not in {"user", "assistant"}:
|
||||
continue
|
||||
if role == "assistant" and not _assistant_entry_visible_in_history(entry):
|
||||
continue
|
||||
content = str(entry.get("content", "")).strip()
|
||||
if not content:
|
||||
continue
|
||||
selected.append({"role": role, "content": content})
|
||||
if role == "user":
|
||||
user_messages += 1
|
||||
if user_messages >= max_turns:
|
||||
break
|
||||
selected.reverse()
|
||||
return selected
|
||||
|
||||
def is_reply_to_assistant(self, session: ChannelSession, message_id: Optional[str]) -> bool:
|
||||
if not message_id:
|
||||
return False
|
||||
for entry in reversed(self._read_transcript_entries(session)):
|
||||
if entry.get("platform_message_id") == message_id:
|
||||
return entry.get("role") == "assistant" and _assistant_entry_visible_in_history(entry)
|
||||
return False
|
||||
|
||||
def list_sessions(self) -> List[ChannelSession]:
|
||||
sessions: List[ChannelSession] = []
|
||||
for metadata_path in sorted(self.sessions_dir.glob("*/session.json")):
|
||||
try:
|
||||
with open(metadata_path, "r", encoding="utf-8") as handle:
|
||||
sessions.append(ChannelSession.from_dict(json.load(handle)))
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to load session metadata %s: %s", metadata_path, exc)
|
||||
return sessions
|
||||
|
||||
def _append_transcript_entry(self, session: ChannelSession, entry: Dict[str, Any]) -> None:
|
||||
transcript_path = Path(session.transcript_path)
|
||||
transcript_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(transcript_path, "a", encoding="utf-8") as handle:
|
||||
handle.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
session.updated_at = _utcnow_iso()
|
||||
self._write_session_metadata(session)
|
||||
|
||||
def _write_session_metadata(self, session: ChannelSession) -> None:
|
||||
with open(session.metadata_path, "w", encoding="utf-8") as handle:
|
||||
json.dump(session.to_dict(), handle, ensure_ascii=False, indent=2)
|
||||
|
||||
def _read_transcript_entries(self, session: ChannelSession) -> List[Dict[str, Any]]:
|
||||
transcript_path = Path(session.transcript_path)
|
||||
if not transcript_path.exists():
|
||||
return []
|
||||
entries: List[Dict[str, Any]] = []
|
||||
with open(transcript_path, "r", encoding="utf-8") as handle:
|
||||
for line in handle:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
entries.append(json.loads(line))
|
||||
except json.JSONDecodeError:
|
||||
logger.warning("Skipping malformed transcript line in %s", transcript_path)
|
||||
return entries
|
||||
|
||||
|
||||
def build_session_key(source: ChannelSource) -> str:
|
||||
parts = [source.platform.value, _sanitize(source.chat_id)]
|
||||
if source.thread_id:
|
||||
parts.append(_sanitize(source.thread_id))
|
||||
return "__".join(part for part in parts if part)
|
||||
|
||||
|
||||
def _sanitize(value: str) -> str:
|
||||
value = re.sub(r"[^a-zA-Z0-9._-]+", "-", str(value).strip())
|
||||
value = value.strip("-._")
|
||||
return value or "unknown"
|
||||
|
||||
|
||||
def _utcnow_iso() -> str:
|
||||
return datetime.now(timezone.utc).isoformat()
|
||||
|
||||
|
||||
def _assistant_entry_visible_in_history(entry: Dict[str, Any]) -> bool:
|
||||
metadata = entry.get("metadata")
|
||||
if not isinstance(metadata, dict):
|
||||
return True
|
||||
return metadata.get("send_success") is not False
|
||||
163
openspace/communication/types.py
Normal file
163
openspace/communication/types.py
Normal file
|
|
@ -0,0 +1,163 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
|
||||
class ChannelPlatform(str, Enum):
|
||||
WHATSAPP = "whatsapp"
|
||||
FEISHU = "feishu"
|
||||
|
||||
|
||||
class AttachmentKind(str, Enum):
|
||||
IMAGE = "image"
|
||||
DOCUMENT = "document"
|
||||
FILE = "file"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ChannelAttachment:
|
||||
kind: AttachmentKind
|
||||
path: str
|
||||
name: str = ""
|
||||
mime_type: str = ""
|
||||
size_bytes: Optional[int] = None
|
||||
source_url: Optional[str] = None
|
||||
metadata: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def to_context_dict(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"kind": self.kind.value,
|
||||
"path": self.path,
|
||||
"name": self.name,
|
||||
"mime_type": self.mime_type,
|
||||
"size_bytes": self.size_bytes,
|
||||
"source_url": self.source_url,
|
||||
"metadata": self.metadata,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ChannelSource:
|
||||
platform: ChannelPlatform
|
||||
chat_id: str
|
||||
chat_type: str = "dm"
|
||||
user_id: Optional[str] = None
|
||||
user_name: Optional[str] = None
|
||||
chat_name: Optional[str] = None
|
||||
thread_id: Optional[str] = None
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"platform": self.platform.value,
|
||||
"chat_id": self.chat_id,
|
||||
"chat_type": self.chat_type,
|
||||
"user_id": self.user_id,
|
||||
"user_name": self.user_name,
|
||||
"chat_name": self.chat_name,
|
||||
"thread_id": self.thread_id,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: Dict[str, Any]) -> "ChannelSource":
|
||||
return cls(
|
||||
platform=ChannelPlatform(str(data["platform"])),
|
||||
chat_id=str(data["chat_id"]),
|
||||
chat_type=str(data.get("chat_type", "dm")),
|
||||
user_id=_optional_str(data.get("user_id")),
|
||||
user_name=_optional_str(data.get("user_name")),
|
||||
chat_name=_optional_str(data.get("chat_name")),
|
||||
thread_id=_optional_str(data.get("thread_id")),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ChannelMessage:
|
||||
source: ChannelSource
|
||||
text: str
|
||||
message_id: str
|
||||
attachments: List[ChannelAttachment] = field(default_factory=list)
|
||||
reply_to_message_id: Optional[str] = None
|
||||
reply_to_text: Optional[str] = None
|
||||
mentions_bot: bool = False
|
||||
metadata: Dict[str, Any] = field(default_factory=dict)
|
||||
received_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
def to_channel_context(self, session_key: str) -> Dict[str, Any]:
|
||||
return {
|
||||
"platform": self.source.platform.value,
|
||||
"chat_id": self.source.chat_id,
|
||||
"chat_type": self.source.chat_type,
|
||||
"chat_name": self.source.chat_name,
|
||||
"thread_id": self.source.thread_id,
|
||||
"user_id": self.source.user_id,
|
||||
"user_name": self.source.user_name,
|
||||
"session_key": session_key,
|
||||
"message_id": self.message_id,
|
||||
"reply_to_message_id": self.reply_to_message_id,
|
||||
"reply_to_text": self.reply_to_text,
|
||||
"attachments": [attachment.to_context_dict() for attachment in self.attachments],
|
||||
}
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ChannelReply:
|
||||
content: str
|
||||
metadata: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SendResult:
|
||||
success: bool
|
||||
message_id: Optional[str] = None
|
||||
error: Optional[str] = None
|
||||
raw_response: Any = None
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ChannelSession:
|
||||
session_key: str
|
||||
source: ChannelSource
|
||||
session_dir: str
|
||||
workspace_dir: str
|
||||
attachments_dir: str
|
||||
transcript_path: str
|
||||
metadata_path: str
|
||||
created_at: str
|
||||
updated_at: str
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"session_key": self.session_key,
|
||||
"source": self.source.to_dict(),
|
||||
"session_dir": self.session_dir,
|
||||
"workspace_dir": self.workspace_dir,
|
||||
"attachments_dir": self.attachments_dir,
|
||||
"transcript_path": self.transcript_path,
|
||||
"metadata_path": self.metadata_path,
|
||||
"created_at": self.created_at,
|
||||
"updated_at": self.updated_at,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: Dict[str, Any]) -> "ChannelSession":
|
||||
return cls(
|
||||
session_key=str(data["session_key"]),
|
||||
source=ChannelSource.from_dict(data["source"]),
|
||||
session_dir=str(data["session_dir"]),
|
||||
workspace_dir=str(data["workspace_dir"]),
|
||||
attachments_dir=str(data["attachments_dir"]),
|
||||
transcript_path=str(data["transcript_path"]),
|
||||
metadata_path=str(data["metadata_path"]),
|
||||
created_at=str(data["created_at"]),
|
||||
updated_at=str(data["updated_at"]),
|
||||
)
|
||||
|
||||
|
||||
def _optional_str(value: Any) -> Optional[str]:
|
||||
if value is None:
|
||||
return None
|
||||
value = str(value).strip()
|
||||
return value or None
|
||||
|
|
@ -33,6 +33,8 @@ Set via `.env`, MCP config `env` block, or system environment.
|
|||
| `OPENSPACE_MODEL` | LLM model | `openrouter/anthropic/claude-sonnet-4.5` |
|
||||
| `OPENSPACE_LLM_API_KEY` | LLM API key (Tier 1 override) | — |
|
||||
| `OPENSPACE_LLM_API_BASE` | LLM API base URL | — |
|
||||
| `OLLAMA_API_BASE` | Local Ollama endpoint for `ollama/*` models | `http://127.0.0.1:11434` |
|
||||
| `OLLAMA_API_KEY` | Placeholder key for Ollama-compatible clients | `ollama` |
|
||||
| `OPENSPACE_LLM_EXTRA_HEADERS` | Extra LLM headers (JSON) | — |
|
||||
| `OPENSPACE_LLM_CONFIG` | Arbitrary litellm kwargs (JSON) | — |
|
||||
| `OPENSPACE_API_KEY` | Cloud API key ([open-space.cloud](https://open-space.cloud)) | — |
|
||||
|
|
@ -91,6 +93,7 @@ Layered system — later files override earlier ones:
|
|||
| `config_mcp.json` | MCP servers OpenSpace connects to as a client |
|
||||
| `config_security.json` | Security policies, blocked commands, sandboxing |
|
||||
| `config_dev.json` | Dev overrides — copy from `config_dev.json.example` (highest priority) |
|
||||
| `config_communication.json` | Communication gateway settings for WhatsApp and Feishu. Use `agent` for per-message OpenSpace execution and `sessions` for queue/history limits. LLM model stays in `openspace/.env`. |
|
||||
|
||||
### Agent config (`config_agents.json`)
|
||||
|
||||
|
|
@ -123,3 +126,41 @@ Layered system — later files override earlier ones:
|
|||
| `blocked_commands` | Platform-specific blacklists (common/linux/darwin/windows) | `rm -rf`, `shutdown`, `dd`, etc. |
|
||||
| `sandbox_enabled` | Enable sandboxing for all operations | `false` |
|
||||
| Per-backend overrides | Shell, MCP, GUI, Web each have independent security policies | Inherit global |
|
||||
|
||||
## 6. Communication Gateway
|
||||
|
||||
The tracked communication config is safe-by-default: loopback-only, channels disabled, and deny-by-default access control. Copy the example config, fill in credentials and `allowed_users`, then explicitly enable the channels you want. The gateway model is not configured here; it inherits `OPENSPACE_MODEL` from `openspace/.env`.
|
||||
|
||||
```bash
|
||||
cp openspace/config/config_communication.json.example openspace/config/config_communication.json
|
||||
```
|
||||
|
||||
Install the Feishu SDK extra when you need Feishu support:
|
||||
|
||||
```bash
|
||||
pip install -e '.[communication]'
|
||||
```
|
||||
|
||||
Start the gateway with either entrypoint:
|
||||
|
||||
```bash
|
||||
openspace communication run --config openspace/config/config_communication.json
|
||||
openspace-gateway --config openspace/config/config_communication.json
|
||||
```
|
||||
|
||||
Check health:
|
||||
|
||||
```bash
|
||||
openspace communication health --config openspace/config/config_communication.json
|
||||
```
|
||||
|
||||
Notes:
|
||||
|
||||
- The tracked `config_communication.json` now stays local-only and deny-by-default. Keep credentials out of git and populate them from a private working copy or environment variables.
|
||||
- Set `server.host` to `0.0.0.0` only when Feishu needs to reach the webhook from outside the machine, and pair that with a populated allowlist plus webhook verification secrets.
|
||||
- Feishu now supports both `webhook` and `websocket` modes. `websocket` matches nanobot's long-connection setup and does not require a public webhook URL.
|
||||
- WhatsApp requires Node.js and npm. The bundled bridge installs its dependencies on first start when `auto_install_dependencies` is enabled.
|
||||
- Set `feishu.bot_open_id` if you want strict group mention gating and automatic bot identity discovery is unavailable in your deployment.
|
||||
- Group chats are gated by `group_policy`. `reply_or_mention` is the default and only accepts messages that mention the bot or reply to a prior assistant message.
|
||||
- `allowed_users` is enforced when `allow_all_users` is `false`. The secure default is deny-by-default until you populate the allowlist.
|
||||
- Attachment caching is limited by `sessions.max_attachment_bytes` and `sessions.max_session_attachment_bytes` to bound disk usage per file and per session.
|
||||
|
|
|
|||
65
openspace/config/config_communication.json.example
Normal file
65
openspace/config/config_communication.json.example
Normal file
|
|
@ -0,0 +1,65 @@
|
|||
{
|
||||
"data_dir": "./logs/communication",
|
||||
"server": {
|
||||
"host": "127.0.0.1",
|
||||
"port": 8765,
|
||||
"health_path": "/health"
|
||||
},
|
||||
"agent": {
|
||||
"max_iterations": 20,
|
||||
"enable_recording": true,
|
||||
"recording_backends": [
|
||||
"shell"
|
||||
],
|
||||
"backend_scope": null,
|
||||
"grounding_config_path": null,
|
||||
"workspace_root": null,
|
||||
"llm_timeout": 120.0
|
||||
},
|
||||
"sessions": {
|
||||
"history_max_turns": 12,
|
||||
"max_parallel_sessions": 2,
|
||||
"idle_ttl_seconds": 900,
|
||||
"per_session_queue_size": 32,
|
||||
"whatsapp_poll_interval_seconds": 1.0,
|
||||
"max_attachment_bytes": 26214400,
|
||||
"max_session_attachment_bytes": 104857600
|
||||
},
|
||||
"whatsapp": {
|
||||
"enabled": false,
|
||||
"allow_all_users": false,
|
||||
"allowed_users": [
|
||||
"15551234567"
|
||||
],
|
||||
"allow_dm": true,
|
||||
"allow_groups": true,
|
||||
"group_policy": "reply_or_mention",
|
||||
"reply_prefix": "OpenSpace\n────────────\n",
|
||||
"bridge": {
|
||||
"host": "127.0.0.1",
|
||||
"port": 3000,
|
||||
"script_path": null,
|
||||
"session_dir": null,
|
||||
"mode": "self-chat",
|
||||
"auto_install_dependencies": true
|
||||
}
|
||||
},
|
||||
"feishu": {
|
||||
"enabled": false,
|
||||
"allow_all_users": false,
|
||||
"allowed_users": [
|
||||
"ou_xxxxxxxxxxxxx"
|
||||
],
|
||||
"allow_dm": true,
|
||||
"allow_groups": true,
|
||||
"group_policy": "reply_or_mention",
|
||||
"app_id": "cli_xxxxxxxxxxxxx",
|
||||
"app_secret": "xxxxxxxxxxxxx",
|
||||
"domain": "feishu",
|
||||
"connection_mode": "webhook",
|
||||
"verification_token": "",
|
||||
"encrypt_key": "",
|
||||
"bot_open_id": "",
|
||||
"webhook_path": "/feishu/webhook"
|
||||
}
|
||||
}
|
||||
|
|
@ -75,6 +75,20 @@ def _pick_first_env(names: tuple[str, ...]) -> str:
|
|||
return ""
|
||||
|
||||
|
||||
def _ensure_local_no_proxy() -> None:
|
||||
required_hosts = ("127.0.0.1", "localhost")
|
||||
for env_name in ("NO_PROXY", "no_proxy"):
|
||||
current = os.environ.get(env_name, "")
|
||||
entries = [entry.strip() for entry in current.split(",") if entry.strip()]
|
||||
updated = False
|
||||
for host in required_hosts:
|
||||
if host not in entries:
|
||||
entries.append(host)
|
||||
updated = True
|
||||
if updated:
|
||||
os.environ[env_name] = ",".join(entries)
|
||||
|
||||
|
||||
def _infer_provider_name(model: str) -> Optional[str]:
|
||||
"""Infer the provider name from a model string using PROVIDER_REGISTRY."""
|
||||
from openspace.host_detection.nanobot import PROVIDER_REGISTRY
|
||||
|
|
@ -222,6 +236,16 @@ def build_llm_kwargs(model: str) -> tuple[str, Dict[str, Any]]:
|
|||
if not resolved_model:
|
||||
resolved_model = _DEFAULT_MODEL
|
||||
|
||||
# Ollama models must use the Ollama-native API base, even when unrelated
|
||||
# OPENSPACE_LLM_* env vars are present for a different provider.
|
||||
if resolved_model.lower().startswith("ollama/"):
|
||||
ollama_base = os.environ.get("OLLAMA_API_BASE", "").strip() or "http://127.0.0.1:11434"
|
||||
_ensure_local_no_proxy()
|
||||
kwargs["api_base"] = ollama_base.rstrip("/")
|
||||
kwargs["api_key"] = os.environ.get("OLLAMA_API_KEY", "").strip() or kwargs.get("api_key") or "ollama"
|
||||
kwargs.pop("extra_headers", None)
|
||||
source = "ollama runtime"
|
||||
|
||||
# Provider-specific adjustments for litellm routing
|
||||
if resolved_model and "minimax" in resolved_model.lower():
|
||||
final_key = kwargs.get("api_key")
|
||||
|
|
|
|||
|
|
@ -104,7 +104,6 @@ class OpenSpace:
|
|||
return
|
||||
|
||||
logger.info("Initializing OpenSpace...")
|
||||
|
||||
try:
|
||||
self._llm_client = LLMClient(
|
||||
model=self.config.llm_model,
|
||||
|
|
@ -313,7 +312,10 @@ class OpenSpace:
|
|||
|
||||
Args:
|
||||
task: Task instruction
|
||||
context: Additional context
|
||||
context: Additional context. Communication callers may pass:
|
||||
- conversation_history: prior user/assistant turns
|
||||
- channel_context: platform/chat metadata and attachments
|
||||
- session_key: stable external session identifier
|
||||
workspace_dir: Working directory
|
||||
max_iterations: Max iterations override
|
||||
task_id: External task ID for recording/logging. If None, generates a random one.
|
||||
|
|
@ -357,9 +359,11 @@ class OpenSpace:
|
|||
|
||||
# Populated inside the try block; used by finally for analysis
|
||||
result: Dict[str, Any] = {}
|
||||
execution_time = 0.0
|
||||
cancelled_exc: Optional[asyncio.CancelledError] = None
|
||||
|
||||
try:
|
||||
execution_context = context or {}
|
||||
execution_context = dict(context) if context else {}
|
||||
execution_context["task_id"] = task_id
|
||||
execution_context["instruction"] = task
|
||||
|
||||
|
|
@ -531,6 +535,20 @@ class OpenSpace:
|
|||
logger.error(f"Task failed: {result.get('error', 'Unknown error')}")
|
||||
logger.info("="*60)
|
||||
|
||||
except asyncio.CancelledError as exc:
|
||||
execution_time = asyncio.get_event_loop().time() - start_time
|
||||
logger.warning("Task execution cancelled")
|
||||
result = {
|
||||
"status": "cancelled",
|
||||
"error": "Task execution cancelled",
|
||||
"response": "",
|
||||
"execution_time": execution_time,
|
||||
"task_id": task_id,
|
||||
"iterations": 0,
|
||||
"tool_executions": [],
|
||||
}
|
||||
cancelled_exc = exc
|
||||
|
||||
except Exception as e:
|
||||
execution_time = asyncio.get_event_loop().time() - start_time
|
||||
tb = traceback.format_exc(limit=10)
|
||||
|
|
@ -569,14 +587,15 @@ class OpenSpace:
|
|||
except Exception as e:
|
||||
logger.warning(f"Failed to stop recording: {e}")
|
||||
|
||||
# Run execution analysis + evolution BEFORE building the return
|
||||
# value, so evolved_skills is populated.
|
||||
await self._maybe_analyze_execution(
|
||||
task_id, recording_dir, result
|
||||
)
|
||||
if cancelled_exc is None:
|
||||
# Run execution analysis + evolution BEFORE building the return
|
||||
# value, so evolved_skills is populated.
|
||||
await self._maybe_analyze_execution(
|
||||
task_id, recording_dir, result
|
||||
)
|
||||
|
||||
# Trigger quality evolution periodically
|
||||
await self._maybe_evolve_quality()
|
||||
# Trigger quality evolution periodically
|
||||
await self._maybe_evolve_quality()
|
||||
|
||||
final_result = {
|
||||
**result,
|
||||
|
|
@ -588,8 +607,10 @@ class OpenSpace:
|
|||
|
||||
self._running = False
|
||||
self._task_done.set()
|
||||
|
||||
return final_result
|
||||
|
||||
if cancelled_exc is not None:
|
||||
raise cancelled_exc
|
||||
return final_result
|
||||
|
||||
# Skills helpers
|
||||
def _init_skill_registry(self) -> Optional[SkillRegistry]:
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ dependencies = [
|
|||
"flask>=3.1.0",
|
||||
"pyautogui>=0.9.54",
|
||||
"pydantic>=2.12.0",
|
||||
"aiohttp>=3.10.0",
|
||||
"requests>=2.32.0",
|
||||
]
|
||||
|
||||
|
|
@ -57,8 +58,15 @@ dev = [
|
|||
"mypy>=1.0.0",
|
||||
]
|
||||
|
||||
communication = [
|
||||
"lark-oapi>=1.4.20",
|
||||
]
|
||||
|
||||
all = [
|
||||
"openspace[macos,linux,windows,dev]",
|
||||
"openspace[macos]; sys_platform == 'darwin'",
|
||||
"openspace[linux]; sys_platform == 'linux'",
|
||||
"openspace[windows]; sys_platform == 'win32'",
|
||||
"openspace[communication,dev]",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
|
|
@ -72,6 +80,7 @@ openspace-mcp = "openspace.mcp_server:run_mcp_server"
|
|||
openspace-download-skill = "openspace.cloud.cli.download_skill:main"
|
||||
openspace-upload-skill = "openspace.cloud.cli.upload_skill:main"
|
||||
openspace-dashboard = "openspace.dashboard_server:main"
|
||||
openspace-gateway = "openspace.communication.gateway:run_main"
|
||||
|
||||
[tool.setuptools]
|
||||
packages = {find = {where = ["."], include = ["openspace*"]}}
|
||||
|
|
@ -80,6 +89,7 @@ packages = {find = {where = ["."], include = ["openspace*"]}}
|
|||
openspace = [
|
||||
"config/*.json",
|
||||
"config/*.json.example",
|
||||
"communication/bridges/whatsapp/*",
|
||||
"local_server/config.json",
|
||||
"local_server/README.md",
|
||||
]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue