diff --git a/.gitignore b/.gitignore index 77dd181..51d57b2 100644 --- a/.gitignore +++ b/.gitignore @@ -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/ diff --git a/openspace/.env.example b/openspace/.env.example index d8e9465..00c4fd6 100644 --- a/openspace/.env.example +++ b/openspace/.env.example @@ -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. diff --git a/openspace/__main__.py b/openspace/__main__.py index 153eebc..ba67d56 100644 --- a/openspace/__main__.py +++ b/openspace/__main__.py @@ -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) diff --git a/openspace/agents/grounding_agent.py b/openspace/agents/grounding_agent.py index d64116c..2ad8a9c 100644 --- a/openspace/agents/grounding_agent.py +++ b/openspace/agents/grounding_agent.py @@ -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. diff --git a/openspace/agents/message_utils.py b/openspace/agents/message_utils.py new file mode 100644 index 0000000..4c3948e --- /dev/null +++ b/openspace/agents/message_utils.py @@ -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) diff --git a/openspace/agents/visual_analyzer.py b/openspace/agents/visual_analyzer.py new file mode 100644 index 0000000..1179672 --- /dev/null +++ b/openspace/agents/visual_analyzer.py @@ -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 diff --git a/openspace/communication/__init__.py b/openspace/communication/__init__.py new file mode 100644 index 0000000..5e5ac01 --- /dev/null +++ b/openspace/communication/__init__.py @@ -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", +] diff --git a/openspace/communication/adapters/__init__.py b/openspace/communication/adapters/__init__.py new file mode 100644 index 0000000..649009a --- /dev/null +++ b/openspace/communication/adapters/__init__.py @@ -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", +] diff --git a/openspace/communication/adapters/base.py b/openspace/communication/adapters/base.py new file mode 100644 index 0000000..fb0fece --- /dev/null +++ b/openspace/communication/adapters/base.py @@ -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 diff --git a/openspace/communication/adapters/feishu.py b/openspace/communication/adapters/feishu.py new file mode 100644 index 0000000..3c11063 --- /dev/null +++ b/openspace/communication/adapters/feishu.py @@ -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 diff --git a/openspace/communication/adapters/whatsapp.py b/openspace/communication/adapters/whatsapp.py new file mode 100644 index 0000000..7218965 --- /dev/null +++ b/openspace/communication/adapters/whatsapp.py @@ -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("+") diff --git a/openspace/communication/attachment_cache.py b/openspace/communication/attachment_cache.py new file mode 100644 index 0000000..4c660fe --- /dev/null +++ b/openspace/communication/attachment_cache.py @@ -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" diff --git a/openspace/communication/bridges/whatsapp/allowlist.js b/openspace/communication/bridges/whatsapp/allowlist.js new file mode 100644 index 0000000..c4a4948 --- /dev/null +++ b/openspace/communication/bridges/whatsapp/allowlist.js @@ -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; +} diff --git a/openspace/communication/bridges/whatsapp/bridge.js b/openspace/communication/bridges/whatsapp/bridge.js new file mode 100644 index 0000000..56e3c2d --- /dev/null +++ b/openspace/communication/bridges/whatsapp/bridge.js @@ -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); + } + }); +} diff --git a/openspace/communication/bridges/whatsapp/package.json b/openspace/communication/bridges/whatsapp/package.json new file mode 100644 index 0000000..0279e20 --- /dev/null +++ b/openspace/communication/bridges/whatsapp/package.json @@ -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" + } +} diff --git a/openspace/communication/config.py b/openspace/communication/config.py new file mode 100644 index 0000000..787839d --- /dev/null +++ b/openspace/communication/config.py @@ -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) diff --git a/openspace/communication/gateway.py b/openspace/communication/gateway.py new file mode 100644 index 0000000..737ae77 --- /dev/null +++ b/openspace/communication/gateway.py @@ -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() diff --git a/openspace/communication/gateway_runtime.py b/openspace/communication/gateway_runtime.py new file mode 100644 index 0000000..cbcdbc4 --- /dev/null +++ b/openspace/communication/gateway_runtime.py @@ -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 diff --git a/openspace/communication/policy.py b/openspace/communication/policy.py new file mode 100644 index 0000000..fb24d0c --- /dev/null +++ b/openspace/communication/policy.py @@ -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 diff --git a/openspace/communication/runtime_manager.py b/openspace/communication/runtime_manager.py new file mode 100644 index 0000000..5148fbe --- /dev/null +++ b/openspace/communication/runtime_manager.py @@ -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() diff --git a/openspace/communication/session_store.py b/openspace/communication/session_store.py new file mode 100644 index 0000000..4c66e39 --- /dev/null +++ b/openspace/communication/session_store.py @@ -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 diff --git a/openspace/communication/types.py b/openspace/communication/types.py new file mode 100644 index 0000000..34ef037 --- /dev/null +++ b/openspace/communication/types.py @@ -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 diff --git a/openspace/config/README.md b/openspace/config/README.md index f2e0875..a65aac3 100644 --- a/openspace/config/README.md +++ b/openspace/config/README.md @@ -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. diff --git a/openspace/config/config_communication.json.example b/openspace/config/config_communication.json.example new file mode 100644 index 0000000..32974d2 --- /dev/null +++ b/openspace/config/config_communication.json.example @@ -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" + } +} diff --git a/openspace/host_detection/resolver.py b/openspace/host_detection/resolver.py index 90b095b..54115e5 100644 --- a/openspace/host_detection/resolver.py +++ b/openspace/host_detection/resolver.py @@ -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") diff --git a/openspace/tool_layer.py b/openspace/tool_layer.py index c2eaca9..7764e1b 100644 --- a/openspace/tool_layer.py +++ b/openspace/tool_layer.py @@ -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]: diff --git a/pyproject.toml b/pyproject.toml index 06c199c..38eba3c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", ]