diff --git a/Dockerfile b/Dockerfile index 5cd7e5e..acde1cb 100644 --- a/Dockerfile +++ b/Dockerfile @@ -4,14 +4,14 @@ RUN apt-get update && apt-get install -y --no-install-recommends curl && rm -rf WORKDIR /app -# Copy application +# Copy project files COPY pyproject.toml . COPY api/ api/ COPY mcp_server/ mcp_server/ COPY config/ config/ -# Install dependencies -RUN pip install --no-cache-dir fastapi "uvicorn[standard]" httpx pydantic openai pyyaml +# Install dependencies from pyproject.toml +RUN pip install --no-cache-dir . # Create persistent data directories RUN mkdir -p /opt/data/skills /opt/data/profiles /opt/data/cases diff --git a/api/routes/health.py b/api/routes/health.py index dde4b60..45139cf 100644 --- a/api/routes/health.py +++ b/api/routes/health.py @@ -8,12 +8,14 @@ from fastapi import APIRouter router = APIRouter() +VERSION = "0.2.0" + @router.get("/api/health") async def health(): return { "status": "ok", "service": "shira-hermes", - "version": "0.1.0", + "version": VERSION, "timestamp": datetime.now(timezone.utc).isoformat(), } diff --git a/api/routes/smart_assistant.py b/api/routes/smart_assistant.py index f65aa3a..3fbb356 100644 --- a/api/routes/smart_assistant.py +++ b/api/routes/smart_assistant.py @@ -19,6 +19,9 @@ logger = logging.getLogger("shira.api") router = APIRouter() +# Hold references to background tasks so they don't get garbage-collected +_background_tasks: set[asyncio.Task] = set() + def _get_api_key() -> str: return os.environ.get("API_KEY", "") @@ -128,7 +131,9 @@ async def smart_assistant_chat(request: Request): # Trigger self-learning in the background (non-blocking) conversation_messages = runner.get_messages() if conversation_messages: - asyncio.create_task(_learn_in_background(conversation_messages)) + task = asyncio.create_task(_learn_in_background(conversation_messages)) + _background_tasks.add(task) + task.add_done_callback(_background_tasks.discard) return { "text": text, diff --git a/api/services/memory.py b/api/services/memory.py index 2004bdf..5b31d6d 100644 --- a/api/services/memory.py +++ b/api/services/memory.py @@ -5,6 +5,7 @@ from __future__ import annotations import asyncio import logging import os +import re import tempfile from pathlib import Path @@ -18,14 +19,21 @@ PROFILE_MAX_CHARS = 1500 CASE_MEMORY_MAX_CHARS = 2500 ENTRY_DELIMITER = "\n§\n" +_SAFE_ID_RE = re.compile(r"^[a-zA-Z0-9_-]+$") + # Per-path locks for thread-safe writes _locks: dict[str, asyncio.Lock] = {} +def _safe_id(value: str) -> str: + """Validate that an ID is safe for use in file paths.""" + if not value or not _SAFE_ID_RE.match(value): + raise ValueError(f"Invalid ID for file path: {value!r}") + return value + + def _get_lock(path: str) -> asyncio.Lock: - if path not in _locks: - _locks[path] = asyncio.Lock() - return _locks[path] + return _locks.setdefault(path, asyncio.Lock()) def _read_file(path: str) -> str: @@ -75,17 +83,17 @@ def _truncate(content: str, max_chars: int) -> str: def load_user_profile(user_id: str) -> str: """Load the per-lawyer profile. Returns empty string if none exists.""" - return _read_file(f"{PROFILES_DIR}/{user_id}.md") + return _read_file(f"{PROFILES_DIR}/{_safe_id(user_id)}.md") def load_case_memory(case_id: str) -> str: """Load the per-case memory. Returns empty string if none exists.""" - return _read_file(f"{CASES_DIR}/{case_id}/memory.md") + return _read_file(f"{CASES_DIR}/{_safe_id(case_id)}/memory.md") async def update_user_profile(user_id: str, action: str, content: str, match: str = "") -> str: """Add, replace, or remove an entry in the lawyer's profile.""" - path = f"{PROFILES_DIR}/{user_id}.md" + path = f"{PROFILES_DIR}/{_safe_id(user_id)}.md" async with _get_lock(path): current = _read_file(path) entries = _entries(current) @@ -129,7 +137,7 @@ async def update_user_profile(user_id: str, action: str, content: str, match: st async def update_case_memory(case_id: str, action: str, content: str, match: str = "") -> str: """Add, replace, or remove an entry in the case memory.""" - path = f"{CASES_DIR}/{case_id}/memory.md" + path = f"{CASES_DIR}/{_safe_id(case_id)}/memory.md" async with _get_lock(path): current = _read_file(path) entries = _entries(current) diff --git a/api/services/prompt_builder.py b/api/services/prompt_builder.py index 03d09f7..bf6fa76 100644 --- a/api/services/prompt_builder.py +++ b/api/services/prompt_builder.py @@ -2,7 +2,8 @@ from __future__ import annotations -from datetime import datetime, timezone, timedelta +from datetime import datetime +from zoneinfo import ZoneInfo VALID_STATUSES = { "New": "חדש", @@ -19,7 +20,7 @@ VALID_STATUSES = { "Rejected": "נדחה", } -ISR_TZ = timezone(timedelta(hours=3)) +ISR_TZ = ZoneInfo("Asia/Jerusalem") TOOL_RULES = ( "TOOL RULES (CRITICAL — you MUST follow these):\n" diff --git a/api/services/skills.py b/api/services/skills.py index ce4f3d0..bfdabca 100644 --- a/api/services/skills.py +++ b/api/services/skills.py @@ -4,6 +4,7 @@ from __future__ import annotations import logging import os +import re from pathlib import Path import yaml @@ -12,6 +13,16 @@ logger = logging.getLogger("shira.skills") SKILLS_DIR = os.environ.get("SHIRA_DATA_DIR", "/opt/data") + "/skills" +_SAFE_NAME_RE = re.compile(r"^[a-zA-Z0-9_-]+$") + + +def _safe_skill_name(name: str) -> str: + """Validate and normalize a skill name for use in file paths.""" + name = name.lower().strip() + if not name or not _SAFE_NAME_RE.match(name): + raise ValueError(f"Invalid skill name: {name!r}") + return name + def _parse_frontmatter(text: str) -> tuple[dict, str]: """Parse YAML frontmatter from a SKILL.md file. Returns (meta, body).""" @@ -53,7 +64,7 @@ def list_skills() -> list[dict]: def view_skill(skill_name: str) -> str | None: """Load the full SKILL.md content for a given skill name.""" - skill_file = Path(SKILLS_DIR) / skill_name / "SKILL.md" + skill_file = Path(SKILLS_DIR) / _safe_skill_name(skill_name) / "SKILL.md" if not skill_file.exists(): return None return skill_file.read_text(encoding="utf-8") @@ -61,10 +72,11 @@ def view_skill(skill_name: str) -> str | None: def save_skill(skill_name: str, content: str) -> None: """Create or overwrite a skill file.""" - skill_dir = Path(SKILLS_DIR) / skill_name + safe_name = _safe_skill_name(skill_name) + skill_dir = Path(SKILLS_DIR) / safe_name skill_dir.mkdir(parents=True, exist_ok=True) (skill_dir / "SKILL.md").write_text(content, encoding="utf-8") - logger.info("Saved skill: %s", skill_name) + logger.info("Saved skill: %s", safe_name) def register_skill_tools(tools: dict) -> None: diff --git a/mcp_server/espocrm_client.py b/mcp_server/espocrm_client.py index 22ed447..40df917 100644 --- a/mcp_server/espocrm_client.py +++ b/mcp_server/espocrm_client.py @@ -8,6 +8,14 @@ import logging logger = logging.getLogger("shira.espocrm") +class EspoCrmApiError(RuntimeError): + """Raised when the EspoCRM API returns a non-success status code.""" + + def __init__(self, status_code: int, text: str): + self.status_code = status_code + super().__init__(f"EspoCRM API error {status_code}: {text}") + + class EspoCrmClient: """Thin async wrapper around the EspoCRM REST API.""" @@ -41,7 +49,7 @@ class EspoCrmClient: if not resp.is_success: text = resp.text logger.error("[crm-api] Error %s: %s", resp.status_code, text) - raise RuntimeError(f"EspoCRM API error {resp.status_code}: {text}") + raise EspoCrmApiError(resp.status_code, text) return resp.json() async def get(self, endpoint: str) -> dict: diff --git a/mcp_server/tools/crm_tools.py b/mcp_server/tools/crm_tools.py index 4786e72..04998bf 100644 --- a/mcp_server/tools/crm_tools.py +++ b/mcp_server/tools/crm_tools.py @@ -7,6 +7,7 @@ from typing import TYPE_CHECKING if TYPE_CHECKING: from mcp_server.espocrm_client import EspoCrmClient +from mcp_server.espocrm_client import EspoCrmApiError from mcp_server.tools._helpers import normalize_date, ok, fail VALID_STATUSES = { @@ -200,13 +201,15 @@ def register_crm_tools( "isActive": True, }) return ok(f'כלל חדש נשמר: "{name}" (ID: {result["id"]})') - except Exception as e: - if "404" in str(e) and user_id: + except EspoCrmApiError as e: + if e.status_code == 404 and user_id: # AssistantRule entity not available — save to user profile instead from api.services.memory import update_user_profile entry = f"כלל [{scope}] {name}: {rule}" return await update_user_profile(user_id, "add", entry) return fail(f"שגיאה בשמירת כלל: {e}") + except Exception as e: + return fail(f"שגיאה בשמירת כלל: {e}") # Register all tools tools["create_task"] = { diff --git a/pyproject.toml b/pyproject.toml index a15daa9..447637e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -8,7 +8,6 @@ dependencies = [ "uvicorn[standard]>=0.32.0", "httpx>=0.27.0", "pydantic>=2.10.0", - "mcp>=1.0.0", "openai>=1.50.0", "pyyaml>=6.0", ]