Spaces:
Sleeping
Sleeping
| from typing import Any | |
| import asyncio | |
| from pathlib import Path | |
| from ace.integrations.mcp.registry import SessionRegistry | |
| from ace.integrations.mcp.models import ( | |
| AskRequest, | |
| AskResponse, | |
| LearnSampleRequest, | |
| LearnSampleResponse, | |
| LearnFeedbackRequest, | |
| LearnFeedbackResponse, | |
| SkillbookGetRequest, | |
| SkillbookGetResponse, | |
| SkillbookSaveRequest, | |
| SkillbookSaveResponse, | |
| SkillbookLoadRequest, | |
| SkillbookLoadResponse, | |
| SkillItem, | |
| ) | |
| from ace.integrations.mcp.config import MCPServerConfig | |
| from ace.integrations.mcp.errors import ( | |
| ACEMCPError, | |
| ForbiddenInSafeModeError, | |
| InternalError, | |
| SaveLoadDisabledError, | |
| TimeoutError as MCPTimeoutError, | |
| ValidationError, | |
| ) | |
| from ace.core.environments import Sample | |
| class MCPHandlers: | |
| def __init__(self, registry: SessionRegistry, config: MCPServerConfig): | |
| self.registry = registry | |
| self.config = config | |
| def _get_session_kwargs(self, config_model) -> tuple[str | None, dict[str, Any]]: | |
| target_model = None | |
| kwargs: dict[str, Any] = {} | |
| if config_model: | |
| target_model = config_model.model # may be None per contract | |
| if config_model.temperature is not None: | |
| kwargs["temperature"] = config_model.temperature | |
| if config_model.max_tokens is not None: | |
| kwargs["max_tokens"] = config_model.max_tokens | |
| return target_model, kwargs | |
| def _enforce_prompt_limit(self, char_count: int, field_name: str) -> None: | |
| if char_count > self.config.max_prompt_chars: | |
| raise ValidationError( | |
| f"{field_name} exceeds max_prompt_chars ({self.config.max_prompt_chars})", | |
| details={ | |
| "field": field_name, | |
| "char_count": char_count, | |
| "max_prompt_chars": self.config.max_prompt_chars, | |
| }, | |
| ) | |
| def _resolve_skillbook_path(self, path: str) -> str: | |
| """Resolve a user-provided path and validate it against skillbook_root. | |
| Returns the resolved absolute path string so callers use the | |
| validated path — not the raw user input — for file operations, | |
| eliminating TOCTOU races with symlinks or ``..`` components. | |
| """ | |
| resolved = str(Path(path).expanduser().resolve()) | |
| if not self.config.skillbook_root: | |
| return resolved | |
| root = Path(self.config.skillbook_root).expanduser().resolve() | |
| try: | |
| Path(resolved).relative_to(root) | |
| except ValueError as exc: | |
| raise ValidationError( | |
| "Path is outside configured skillbook_root", | |
| details={ | |
| "path": resolved, | |
| "skillbook_root": str(root), | |
| }, | |
| ) from exc | |
| return resolved | |
| async def handle_ask(self, request: AskRequest) -> AskResponse: | |
| self._enforce_prompt_limit(len(request.question) + len(request.context), "ask") | |
| target_model, kwargs = self._get_session_kwargs(request.session_config) | |
| session = await self.registry.get_or_create( | |
| request.session_id, model=target_model, **kwargs | |
| ) | |
| async with session.lock: | |
| try: | |
| answer = await asyncio.to_thread( | |
| session.runner.ask, request.question, request.context | |
| ) | |
| skill_count = len(session.runner.skillbook.skills()) | |
| return AskResponse( | |
| session_id=request.session_id, | |
| answer=str(answer), | |
| skill_count=skill_count, | |
| ) | |
| except ACEMCPError: | |
| raise | |
| except Exception as e: | |
| raise InternalError(str(e)) | |
| async def handle_skillbook_get( | |
| self, request: SkillbookGetRequest | |
| ) -> SkillbookGetResponse: | |
| session = await self.registry.get(request.session_id) | |
| async with session.lock: | |
| try: | |
| skillbook = session.runner.skillbook | |
| skills = skillbook.skills(include_invalid=request.include_invalid) | |
| limited_skills: list[SkillItem] = [] | |
| for s in skills: | |
| content = getattr(s, "insight", None) or getattr(s, "issue", None) | |
| limited_skills.append( | |
| SkillItem( | |
| id=getattr(s, "id", str(len(limited_skills))), | |
| content=content if content is not None else str(s), | |
| topic=getattr(s, "section", None), | |
| helpful=getattr(s, "helpful_count", None), | |
| harmful=getattr(s, "harmful_count", None), | |
| neutral=getattr(s, "neutral_count", None), | |
| ) | |
| ) | |
| limited_skills = limited_skills[: request.limit] | |
| stats = skillbook.stats() | |
| return SkillbookGetResponse( | |
| session_id=request.session_id, stats=stats, skills=limited_skills | |
| ) | |
| except ACEMCPError: | |
| raise | |
| except Exception as e: | |
| raise InternalError(str(e)) | |
| async def handle_learn_sample( | |
| self, request: LearnSampleRequest | |
| ) -> LearnSampleResponse: | |
| if self.config.safe_mode: | |
| raise ForbiddenInSafeModeError("ace.learn.sample") | |
| if len(request.samples) > self.config.max_samples_per_call: | |
| raise ValidationError( | |
| f"samples exceeds max_samples_per_call ({self.config.max_samples_per_call})", | |
| details={ | |
| "sample_count": len(request.samples), | |
| "max_samples_per_call": self.config.max_samples_per_call, | |
| }, | |
| ) | |
| for idx, s in enumerate(request.samples): | |
| self._enforce_prompt_limit( | |
| len(s.question) + len(s.context), | |
| f"samples[{idx}]", | |
| ) | |
| target_model, kwargs = self._get_session_kwargs(request.session_config) | |
| session = await self.registry.get_or_create( | |
| request.session_id, model=target_model, **kwargs | |
| ) | |
| async with session.lock: | |
| try: | |
| samples = [] | |
| for s in request.samples: | |
| samples.append( | |
| Sample( | |
| question=s.question, | |
| context=s.context, | |
| ground_truth=s.ground_truth, | |
| metadata=s.metadata or {}, | |
| ) | |
| ) | |
| count_before = len(session.runner.skillbook.skills()) | |
| results = await asyncio.wait_for( | |
| asyncio.to_thread( | |
| session.runner.learn, | |
| samples, | |
| None, | |
| request.epochs, | |
| ), | |
| timeout=self.config.learn_timeout_seconds, | |
| ) | |
| failed = sum(1 for r in results if r.error is not None) | |
| count_after = len(session.runner.skillbook.skills()) | |
| return LearnSampleResponse( | |
| session_id=request.session_id, | |
| processed=len(samples) - failed, | |
| failed=failed, | |
| skill_count_before=count_before, | |
| skill_count_after=count_after, | |
| new_skill_count=max(0, count_after - count_before), | |
| ) | |
| except ACEMCPError: | |
| raise | |
| except asyncio.TimeoutError: | |
| raise MCPTimeoutError( | |
| f"learn.sample timed out after {self.config.learn_timeout_seconds}s" | |
| ) | |
| except Exception as e: | |
| raise InternalError(str(e)) | |
| async def handle_learn_feedback( | |
| self, request: LearnFeedbackRequest | |
| ) -> LearnFeedbackResponse: | |
| if self.config.safe_mode: | |
| raise ForbiddenInSafeModeError("ace.learn.feedback") | |
| self._enforce_prompt_limit( | |
| len(request.question) | |
| + len(request.context) | |
| + len(request.answer) | |
| + len(request.feedback) | |
| + len(request.ground_truth or ""), | |
| "learn.feedback", | |
| ) | |
| target_model, kwargs = self._get_session_kwargs(request.session_config) | |
| session = await self.registry.get_or_create( | |
| request.session_id, model=target_model, **kwargs | |
| ) | |
| async with session.lock: | |
| try: | |
| count_before = len(session.runner.skillbook.skills()) | |
| # Prefer the direct feedback path when a prior ask exists; | |
| # fall back to learn_from_traces for standalone feedback. | |
| timeout = self.config.learn_timeout_seconds | |
| learned = await asyncio.wait_for( | |
| asyncio.to_thread( | |
| session.runner.learn_from_feedback, | |
| request.feedback, | |
| request.ground_truth or None, | |
| ), | |
| timeout=timeout, | |
| ) | |
| if not learned: | |
| # No prior ask interaction — build a trace and learn | |
| trace: dict[str, object] = { | |
| "question": request.question, | |
| "context": request.context, | |
| "answer": request.answer, | |
| "skill_ids": [], | |
| "feedback": request.feedback, | |
| "ground_truth": request.ground_truth, | |
| } | |
| await asyncio.wait_for( | |
| asyncio.to_thread(session.runner.learn_from_traces, [trace]), | |
| timeout=timeout, | |
| ) | |
| count_after = len(session.runner.skillbook.skills()) | |
| return LearnFeedbackResponse( | |
| session_id=request.session_id, | |
| learned=True, | |
| skill_count_before=count_before, | |
| skill_count_after=count_after, | |
| new_skill_count=max(0, count_after - count_before), | |
| ) | |
| except ACEMCPError: | |
| raise | |
| except asyncio.TimeoutError: | |
| raise MCPTimeoutError( | |
| f"learn.feedback timed out after {self.config.learn_timeout_seconds}s" | |
| ) | |
| except Exception as e: | |
| raise InternalError(str(e)) | |
| async def handle_skillbook_save( | |
| self, request: SkillbookSaveRequest | |
| ) -> SkillbookSaveResponse: | |
| if self.config.safe_mode: | |
| raise ForbiddenInSafeModeError("ace.skillbook.save") | |
| if not self.config.allow_save_load: | |
| raise SaveLoadDisabledError("ace.skillbook.save") | |
| resolved = self._resolve_skillbook_path(request.path) | |
| session = await self.registry.get(request.session_id) | |
| async with session.lock: | |
| try: | |
| await asyncio.to_thread(session.runner.save, resolved) | |
| skill_count = len(session.runner.skillbook.skills()) | |
| return SkillbookSaveResponse( | |
| session_id=request.session_id, | |
| path=resolved, | |
| saved_skill_count=skill_count, | |
| ) | |
| except ACEMCPError: | |
| raise | |
| except Exception as e: | |
| raise InternalError(str(e)) | |
| async def handle_skillbook_load( | |
| self, request: SkillbookLoadRequest | |
| ) -> SkillbookLoadResponse: | |
| if self.config.safe_mode: | |
| raise ForbiddenInSafeModeError("ace.skillbook.load") | |
| if not self.config.allow_save_load: | |
| raise SaveLoadDisabledError("ace.skillbook.load") | |
| resolved = self._resolve_skillbook_path(request.path) | |
| session = await self.registry.get_or_create(request.session_id) | |
| async with session.lock: | |
| try: | |
| await asyncio.to_thread(session.runner.load, resolved) | |
| skill_count = len(session.runner.skillbook.skills()) | |
| return SkillbookLoadResponse( | |
| session_id=request.session_id, | |
| path=resolved, | |
| skill_count=skill_count, | |
| ) | |
| except ACEMCPError: | |
| raise | |
| except Exception as e: | |
| raise InternalError(str(e)) | |