diff --git a/flocks/session/goal.py b/flocks/session/goal.py index 29c1b5b80..0e52adf11 100644 --- a/flocks/session/goal.py +++ b/flocks/session/goal.py @@ -11,6 +11,12 @@ from flocks.provider.options import build_provider_options from flocks.provider.provider import ChatMessage, Provider +from flocks.session.llm_hook_utils import ( + apply_hook_request_output, + restore_text_with_replacements, + serialize_chat_message, + stream_text_replacements_from_hook_output, +) from flocks.storage.storage import Storage from flocks.utils.log import Log @@ -174,6 +180,7 @@ async def judge_goal_with_model( *, provider_id: str, model_id: str, + session_id: Optional[str] = None, initial_clarification: Optional[GoalClarification] = None, ) -> tuple[GoalVerdict, str]: """Hermes-style model judge using the active session provider/model.""" @@ -183,26 +190,84 @@ async def judge_goal_with_model( provider_options = build_provider_options(provider_id, model_id) provider_options.pop("max_tokens", None) + messages = [ + ChatMessage(role="system", content=_MODEL_JUDGE_SYSTEM_PROMPT), + ChatMessage( + role="user", + content=( + f"Goal:\n{_format_goal_context(objective, initial_clarification)}\n\n" + "Latest assistant final response (truncated to the last 4KB):\n" + f"{_judge_input(last_response)}" + ), + ), + ] + replacements: list[tuple[str, str]] = [] + if session_id: + from flocks.hooks.pipeline import HookPipeline, HookStage + + hook_metadata = { + "sessionID": session_id, + "agent": "goal_judge", + "model": { + "providerID": provider_id, + "modelID": model_id, + }, + "purpose": "goal_judge", + } + try: + llm_before_enabled = await HookPipeline.has_stage_handlers( + HookStage.LLM_BEFORE, + hook_metadata, + ) + except Exception as exc: + log.error("goal.model_judge.hook_stage_probe_failed", { + "session_id": session_id, + "provider_id": provider_id, + "model_id": model_id, + "error": str(exc), + }) + raise RuntimeError("Goal judge before-hook stage probe failed; request was not sent") from exc + + if llm_before_enabled: + try: + hook_ctx = await HookPipeline.run_llm_before({ + **hook_metadata, + "request": { + "messageCount": len(messages), + "messages": [serialize_chat_message(message) for message in messages], + "toolCount": 0, + "tools": [], + "providerOptions": dict(provider_options), + "providerToolsEnabled": False, + }, + }) + hook_output = getattr(hook_ctx, "output", None) or {} + replacements = stream_text_replacements_from_hook_output(hook_output) + messages, provider_options = apply_hook_request_output( + messages, + provider_options, + hook_output, + ) + provider_options.pop("max_tokens", None) + except Exception as exc: + log.error("goal.model_judge.llm_before_failed", { + "session_id": session_id, + "provider_id": provider_id, + "model_id": model_id, + "error": str(exc), + }) + raise RuntimeError("Goal judge before-hook failed; request was not sent") from exc response = await provider.chat( model_id=model_id, - messages=[ - ChatMessage(role="system", content=_MODEL_JUDGE_SYSTEM_PROMPT), - ChatMessage( - role="user", - content=( - f"Goal:\n{_format_goal_context(objective, initial_clarification)}\n\n" - "Latest assistant final response (truncated to the last 4KB):\n" - f"{_judge_input(last_response)}" - ), - ), - ], + messages=messages, **provider_options, max_tokens=JUDGE_MAX_TOKENS, temperature=0, ) - payload = _extract_json_object(response.content) + response_content = restore_text_with_replacements(response.content, replacements) + payload = _extract_json_object(response_content) verdict = str(payload.get("verdict") or "").strip().lower() reason = _trim_reason(str(payload.get("reason") or "")) if verdict not in {"complete", "blocked", "waiting", "continue"}: @@ -352,6 +417,7 @@ async def evaluate_after_turn( last_response, provider_id=provider_id, model_id=model_id, + session_id=session_id, initial_clarification=state.initial_clarification, ) except Exception as exc: diff --git a/flocks/session/lifecycle/compaction/compaction.py b/flocks/session/lifecycle/compaction/compaction.py index 7e614b6d8..dcc6aacea 100644 --- a/flocks/session/lifecycle/compaction/compaction.py +++ b/flocks/session/lifecycle/compaction/compaction.py @@ -1190,6 +1190,7 @@ async def process( focus_instruction=focus_instruction, previous_summary=previous_summary, chat_messages=chat_messages, + session_id=session_id, ) except RuntimeError as _e: # No provider configured — long cooldown (hermes: 600s) diff --git a/flocks/session/lifecycle/compaction/summary.py b/flocks/session/lifecycle/compaction/summary.py index 97557eb5e..394efe462 100644 --- a/flocks/session/lifecycle/compaction/summary.py +++ b/flocks/session/lifecycle/compaction/summary.py @@ -13,8 +13,16 @@ # here in the future). ProgressCallback = Callable[[str, Dict[str, Any]], Awaitable[None]] +from flocks.provider.provider import ChatMessage from flocks.utils.log import Log from flocks.session.prompt import SessionPrompt +from flocks.session.llm_hook_utils import ( + apply_hook_request_output, + restore_text_with_replacements, + serialize_chat_message, + stream_text_replacements_from_hook_output, + strip_think_blocks, +) from .models import DEFAULT_COMPACTION_PROMPT_WITH_PREVIOUS log = Log.create(service="session.compaction.summarization") @@ -159,16 +167,72 @@ async def _llm_chat_with_timeout( messages: list, max_tokens: int, timeout: int = COMPACTION_TIMEOUT_SECONDS, + session_id: Optional[str] = None, + purpose: str = "compaction_summary", ) -> Any: """Call provider_client.chat with a timeout guard.""" - return await asyncio.wait_for( + provider_options: Dict[str, Any] = {"max_tokens": max_tokens} + replacements: list[tuple[str, str]] = [] + if session_id: + try: + from flocks.hooks.pipeline import HookPipeline, HookStage + + provider_id = ( + getattr(provider_client, "provider_id", None) + or getattr(provider_client, "id", None) + or provider_client.__class__.__name__ + ) + llm_hook_metadata = { + "sessionID": session_id, + "agent": "session.compaction", + "step": None, + "model": { + "providerID": provider_id, + "modelID": model_id, + }, + "purpose": purpose, + } + if await HookPipeline.has_stage_handlers(HookStage.LLM_BEFORE, llm_hook_metadata): + llm_before_ctx = await HookPipeline.run_llm_before({ + **llm_hook_metadata, + "request": { + "messageCount": len(messages), + "messages": [serialize_chat_message(message) for message in messages], + "toolCount": 0, + "tools": [], + "providerOptions": dict(provider_options), + "providerToolsEnabled": False, + }, + }) + hook_output = getattr(llm_before_ctx, "output", None) or {} + messages, provider_options = apply_hook_request_output( + messages, + provider_options, + hook_output, + ) + replacements = stream_text_replacements_from_hook_output(hook_output) + except Exception as hook_err: + log.error("compaction.llm_before_hook.error", { + "session_id": session_id, + "purpose": purpose, + "error": str(hook_err), + }) + raise RuntimeError("compaction llm_before hook failed; request was not sent") from hook_err + + provider_options.setdefault("max_tokens", max_tokens) + response = await asyncio.wait_for( provider_client.chat( model_id=model_id, messages=messages, - max_tokens=max_tokens, + **provider_options, ), timeout=timeout, ) + if replacements and response is not None and isinstance(getattr(response, "content", None), str): + response.content = restore_text_with_replacements(response.content, replacements) + if response is not None and isinstance(getattr(response, "content", None), str): + response.content = strip_think_blocks(response.content) + return response async def summarize_single_pass( @@ -181,6 +245,7 @@ async def summarize_single_pass( focus_instruction: Optional[str] = None, previous_summary: Optional[str] = None, chat_messages: Optional[list] = None, + session_id: Optional[str] = None, ) -> Optional[str]: """Generate summary in a single LLM call. @@ -201,8 +266,6 @@ async def summarize_single_pass( as "merge new turns into the prior summary" rather than compressing from scratch. """ - from flocks.provider.provider import ChatMessage - if chat_messages: # Per-message truncation path (hermes-style): every turn contributes # a capped fragment (head + tail per message), so early decisions @@ -236,6 +299,8 @@ async def summarize_single_pass( model_id=model_id, messages=[ChatMessage(role="user", content=request)], max_tokens=max_tokens, + session_id=session_id, + purpose="compaction_summary", ) except asyncio.TimeoutError: log.error("compaction.single_pass.timeout", { @@ -452,6 +517,8 @@ async def summarize_chunked_iterative( messages=[ChatMessage(role="user", content=chunk_prompt)], max_tokens=chunk_max_tokens, timeout=COMPACTION_TIMEOUT_SECONDS, + session_id=session_id, + purpose="compaction_summary_chunk", ) duration_ms = (time.perf_counter() - started) * 1000 if resp and resp.content: diff --git a/flocks/session/lifecycle/title.py b/flocks/session/lifecycle/title.py index 6378dd04f..f9bdf5680 100644 --- a/flocks/session/lifecycle/title.py +++ b/flocks/session/lifecycle/title.py @@ -12,6 +12,12 @@ from flocks.utils.log import Log from flocks.provider.provider import ChatMessage +from flocks.session.llm_hook_utils import ( + apply_hook_request_output, + restore_text_with_replacements, + serialize_chat_message, + stream_text_replacements_from_hook_output, +) EventPublishCallback = Optional[Callable[[str, Dict[str, Any]], Awaitable[None]]] @@ -132,16 +138,61 @@ async def generate_title_after_first_message( # Send PROMPT_TITLE as system instruction, user question as user message title = "" try: + title_messages = [ + ChatMessage(role="system", content=_CANONICAL_TITLE_PROMPT), + ChatMessage(role="user", content=question), + ] + provider_options: Dict[str, Any] = {"max_tokens": 50} + replacements: list[tuple[str, str]] = [] + try: + from flocks.hooks.pipeline import HookPipeline, HookStage + + llm_hook_metadata = { + "sessionID": session_id, + "agent": "session.title", + "step": None, + "model": { + "providerID": provider_id, + "modelID": model_id, + }, + "purpose": "title_generation", + } + if await HookPipeline.has_stage_handlers(HookStage.LLM_BEFORE, llm_hook_metadata): + llm_before_ctx = await HookPipeline.run_llm_before({ + **llm_hook_metadata, + "request": { + "messageCount": len(title_messages), + "messages": [serialize_chat_message(message) for message in title_messages], + "toolCount": 0, + "tools": [], + "providerOptions": dict(provider_options), + "providerToolsEnabled": False, + }, + }) + hook_output = getattr(llm_before_ctx, "output", None) or {} + title_messages, provider_options = apply_hook_request_output( + title_messages, + provider_options, + hook_output, + ) + replacements = stream_text_replacements_from_hook_output(hook_output) + except Exception as hook_err: + log.error("title.llm_before_hook.error", { + "session_id": session_id, + "error": str(hook_err), + }) + raise RuntimeError("title llm_before hook failed; request was not sent") from hook_err + + provider_options.setdefault("max_tokens", 50) async for chunk in provider.chat_stream( model_id, - [ - ChatMessage(role="system", content=_CANONICAL_TITLE_PROMPT), - ChatMessage(role="user", content=question), - ], - max_tokens=50, + title_messages, + **provider_options, ): if hasattr(chunk, 'delta') and chunk.delta: title += chunk.delta + if replacements: + title = restore_text_with_replacements(title, replacements) except Exception as llm_err: log.warn("title.llm_failed", { "session_id": session_id, diff --git a/flocks/session/llm_hook_utils.py b/flocks/session/llm_hook_utils.py new file mode 100644 index 000000000..1112dc31e --- /dev/null +++ b/flocks/session/llm_hook_utils.py @@ -0,0 +1,140 @@ +"""Shared helpers for LLM hook request/response payload handling.""" + +from __future__ import annotations + +import re +from typing import Any, Dict, List, Tuple + +from flocks.provider.provider import ChatMessage + + +class StreamingTextReplacementBuffer: + """Incrementally replace streamed placeholders without leaking partial tokens.""" + + def __init__(self, replacements: List[Tuple[str, str]]): + self._replacements = [ + (pattern, value) + for pattern, value in sorted(replacements, key=lambda item: len(item[0]), reverse=True) + if pattern + ] + self._buffer = "" + self._prefixes: set[str] = set() + self._max_pattern_len = 0 + for pattern, _value in self._replacements: + self._max_pattern_len = max(self._max_pattern_len, len(pattern)) + for index in range(1, len(pattern)): + self._prefixes.add(pattern[:index]) + + @property + def enabled(self) -> bool: + return bool(self._replacements) + + def feed(self, text: str) -> str: + if not self.enabled or not text: + return text + self._buffer += text + keep = self._pending_suffix_length(self._buffer) + if keep: + flush_text = self._buffer[:-keep] + self._buffer = self._buffer[-keep:] + else: + flush_text = self._buffer + self._buffer = "" + return restore_text_with_replacements(flush_text, self._replacements) + + def flush(self) -> str: + if not self.enabled or not self._buffer: + return "" + flush_text = self._buffer + self._buffer = "" + return restore_text_with_replacements(flush_text, self._replacements) + + def _pending_suffix_length(self, text: str) -> int: + max_keep = min(len(text), max(self._max_pattern_len - 1, 0)) + for length in range(max_keep, 0, -1): + if text[-length:] in self._prefixes: + return length + return 0 + + +def serialize_chat_message(message: ChatMessage) -> Dict[str, Any]: + payload = message.model_dump(exclude_none=True) + if not payload.get("custom_settings"): + payload.pop("custom_settings", None) + return payload + + +def stream_text_replacements_from_hook_output(output: Dict[str, Any]) -> List[Tuple[str, str]]: + redaction = output.get("redaction") if isinstance(output, dict) else None + raw_items = redaction.get("streamTextReplacements") if isinstance(redaction, dict) else None + if not isinstance(raw_items, list): + return [] + + replacements: List[Tuple[str, str]] = [] + for item in raw_items: + if not isinstance(item, dict): + continue + placeholder = item.get("placeholder") + value = item.get("value") + if isinstance(placeholder, str) and isinstance(value, str): + replacements.append((placeholder, value)) + return replacements + + +def restore_text_with_replacements(text: str, replacements: List[Tuple[str, str]]) -> str: + restored = text + for pattern, value in sorted(replacements, key=lambda item: len(item[0]), reverse=True): + if pattern: + restored = restored.replace(pattern, value) + return restored + + +def restore_value_with_replacements(value: Any, replacements: List[Tuple[str, str]]) -> Any: + if isinstance(value, str): + return restore_text_with_replacements(value, replacements) + if isinstance(value, list): + return [restore_value_with_replacements(item, replacements) for item in value] + if isinstance(value, dict): + return { + key: restore_value_with_replacements(item, replacements) + for key, item in value.items() + } + return value + + +_THINK_BLOCK_RE = re.compile(r".*?", re.DOTALL | re.IGNORECASE) +_THINK_TAG_RE = re.compile(r"", re.IGNORECASE) + + +def strip_think_blocks(text: str) -> str: + if not isinstance(text, str) or " tuple[List[ChatMessage], Dict[str, Any]]: + updated_request = output.get("request") if isinstance(output, dict) else None + if not isinstance(updated_request, dict): + return messages, provider_options + + updated_messages = messages + raw_messages = updated_request.get("messages") + if isinstance(raw_messages, list): + updated_messages = [ + message + if isinstance(message, ChatMessage) + else ChatMessage.model_validate(message) + for message in raw_messages + ] + + updated_provider_options = provider_options + raw_provider_options = updated_request.get("providerOptions") + if isinstance(raw_provider_options, dict): + updated_provider_options = dict(raw_provider_options) + + return updated_messages, updated_provider_options diff --git a/flocks/session/runner.py b/flocks/session/runner.py index e9c10b5ae..86613b5f3 100644 --- a/flocks/session/runner.py +++ b/flocks/session/runner.py @@ -33,6 +33,13 @@ ) from flocks.session.lifecycle.retry import CONNECTION_ERROR_DISPLAY_MESSAGE, SessionRetry from flocks.session.lifecycle.compaction import SessionCompaction, CompactionPolicy +from flocks.session.llm_hook_utils import ( + StreamingTextReplacementBuffer, + apply_hook_request_output, + restore_value_with_replacements, + serialize_chat_message, + stream_text_replacements_from_hook_output, +) from flocks.session.streaming.stream_processor import StreamProcessor from flocks.session.streaming.stream_events import ( StartEvent, @@ -2471,12 +2478,6 @@ async def _call_llm( Uses StreamProcessor to handle events and execute tools synchronously. Ported from Flocks' SessionProcessor.process() behavior. """ - def _serialize_message(message: ChatMessage) -> Dict[str, Any]: - payload = message.model_dump(exclude_none=True) - if not payload.get("custom_settings"): - payload.pop("custom_settings", None) - return payload - def _build_llm_response_payload( *, content: str, @@ -2541,64 +2542,8 @@ def _build_llm_response_payload( reasoning_id_counter = 0 stream_finish_reason: Optional[str] = None - # -- Observability: create trace & generation scopes (safe no-op when - # Langfuse is unconfigured). All observability calls are wrapped in - # try/except so they never break the core session flow. trace_ctx = None generation_ctx = None - if langfuse_is_active(): - try: - input_preview = [] - for _msg in messages[-12:]: - _mc = _msg.content or "" - input_preview.append( - {"role": _msg.role, "chars": len(_mc), "preview": _mc[:240]} - ) - - trace_tags = [ - f"session:{self.session.id}", - f"step:{self._step}", - f"session_step:{self.session.id}:{self._step}", - f"agent:{agent.name}", - f"provider:{self.provider_id}", - ] - trace_ctx = trace_scope( - name="SessionRunner.step", - session_id=self.session.id, - tags=trace_tags, - input={ - "step": self._step, - "message_count": len(messages), - "tool_count": len(tools), - "last_user_preview": next( - ((m.content or "")[:280] for m in reversed(messages) if m.role == "user"), - "", - ), - }, - metadata={ - "provider_id": self.provider_id, - "model_id": self.model_id, - "agent": agent.name, - "workspace": self.session.directory, - }, - ) - generation_ctx = generation_scope( - parent=trace_ctx.observation, - name="LLM.generate", - model=self.model_id, - input=input_preview, - metadata={ - "provider_id": self.provider_id, - "session_id": self.session.id, - "step": self._step, - "tool_names": [t.get("function", {}).get("name", "") for t in tools][:50], - }, - ) - processor._langfuse_generation = generation_ctx.observation - except Exception as exc: - log.debug("runner.observability.init_failed", {"error": str(exc)}) - trace_ctx = None - generation_ctx = None # Validate messages - ensure we have at least one non-system message non_system_messages = [m for m in messages if m.role != "system"] @@ -2646,7 +2591,24 @@ def _build_llm_response_payload( } llm_before_enabled = False llm_after_enabled = False + replacements: list[tuple[str, str]] = [] + stream_text_rewriter: Optional[StreamingTextReplacementBuffer] = None + stream_reasoning_rewriter: Optional[StreamingTextReplacementBuffer] = None self._llm_call_aborted = False + + async def _flush_reasoning_rewriter() -> None: + if stream_reasoning_rewriter is None or not hasattr(self, '_current_reasoning_id'): + return + trailing_reasoning = stream_reasoning_rewriter.flush() + if not trailing_reasoning: + return + reasoning_metadata = getattr(self, '_current_reasoning_metadata', {}) or {} + await processor.process_event(ReasoningDeltaEvent( + id=self._current_reasoning_id, + text=trailing_reasoning, + metadata=reasoning_metadata, + )) + try: llm_before_enabled = await HookPipeline.has_stage_handlers( HookStage.LLM_BEFORE, @@ -2657,14 +2619,15 @@ def _build_llm_response_payload( llm_hook_metadata, ) except Exception as exc: - log.debug("runner.hook.stage_probe.error", {"error": str(exc)}) + log.error("runner.hook.stage_probe.error", {"error": str(exc)}) + raise RuntimeError("LLM hook stage probe failed; request was not sent") from exc if llm_before_enabled: llm_before_hook_input = { **llm_hook_metadata, "request": { "messageCount": len(messages), - "messages": [_serialize_message(message) for message in messages], + "messages": [serialize_chat_message(message) for message in messages], "toolCount": len(tools), "tools": copy.deepcopy(tools), "providerOptions": dict(provider_options), @@ -2673,7 +2636,23 @@ def _build_llm_response_payload( } try: hook_started_at = time.perf_counter() - await HookPipeline.run_llm_before(llm_before_hook_input) + llm_before_ctx = await HookPipeline.run_llm_before(llm_before_hook_input) + hook_output = getattr(llm_before_ctx, "output", None) or {} + replacements = stream_text_replacements_from_hook_output(hook_output) + if replacements: + stream_text_rewriter = StreamingTextReplacementBuffer(replacements) + stream_reasoning_rewriter = StreamingTextReplacementBuffer(replacements) + updated_request = hook_output.get("request") + if isinstance(updated_request, dict): + messages, provider_options = apply_hook_request_output( + messages, + provider_options, + hook_output, + ) + updated_tools = updated_request.get("tools") + if isinstance(updated_tools, list): + tools = copy.deepcopy(updated_tools) + provider_tools = None if self._should_use_text_tool_call_mode() else (tools if tools else None) self._log_perf( "runner.hook.llm_before.complete", hook_started_at, @@ -2681,7 +2660,66 @@ def _build_llm_response_payload( tool_count=len(tools), ) except Exception as exc: - log.debug("runner.hook.llm_before.error", {"error": str(exc)}) + log.error("runner.hook.llm_before.error", {"error": str(exc)}) + raise RuntimeError("LLM before-hook failed; request was not sent") from exc + + # -- Observability: create trace & generation scopes after llm_before, + # so previews use the same redacted messages that will be sent to the provider. + # All observability calls are wrapped in try/except so they never break + # the core session flow. + if langfuse_is_active(): + try: + input_preview = [] + for _msg in messages[-12:]: + _mc = _msg.content or "" + input_preview.append( + {"role": _msg.role, "chars": len(_mc), "preview": _mc[:240]} + ) + + trace_tags = [ + f"session:{self.session.id}", + f"step:{self._step}", + f"session_step:{self.session.id}:{self._step}", + f"agent:{agent.name}", + f"provider:{self.provider_id}", + ] + trace_ctx = trace_scope( + name="SessionRunner.step", + session_id=self.session.id, + tags=trace_tags, + input={ + "step": self._step, + "message_count": len(messages), + "tool_count": len(tools), + "last_user_preview": next( + ((m.content or "")[:280] for m in reversed(messages) if m.role == "user"), + "", + ), + }, + metadata={ + "provider_id": self.provider_id, + "model_id": self.model_id, + "agent": agent.name, + "workspace": self.session.directory, + }, + ) + generation_ctx = generation_scope( + parent=trace_ctx.observation, + name="LLM.generate", + model=self.model_id, + input=input_preview, + metadata={ + "provider_id": self.provider_id, + "session_id": self.session.id, + "step": self._step, + "tool_names": [t.get("function", {}).get("name", "") for t in tools][:50], + }, + ) + processor._langfuse_generation = generation_ctx.observation + except Exception as exc: + log.debug("runner.observability.init_failed", {"error": str(exc)}) + trace_ctx = None + generation_ctx = None llm_call_started_at = time.perf_counter() first_chunk_logged = False @@ -2733,23 +2771,29 @@ def _build_llm_response_payload( # reasoning text. event_type = getattr(chunk, 'event_type', None) chunk_metadata = getattr(chunk, 'metadata', None) or {} + display_chunk_metadata = ( + restore_value_with_replacements(chunk_metadata, replacements) + if replacements + else chunk_metadata + ) reasoning_event_types = {"reasoning", "reasoning-start", "reasoning-end"} - if hasattr(self, '_current_reasoning_id') and chunk_metadata: + if hasattr(self, '_current_reasoning_id') and display_chunk_metadata: current_metadata = getattr(self, '_current_reasoning_metadata', {}) or {} - current_metadata.update(chunk_metadata) + current_metadata.update(display_chunk_metadata) self._current_reasoning_metadata = current_metadata if event_type == "reasoning-start" and not hasattr(self, '_current_reasoning_id'): reasoning_id_counter += 1 self._current_reasoning_id = f"reasoning-{reasoning_id_counter}" - self._current_reasoning_metadata = dict(chunk_metadata) + self._current_reasoning_metadata = dict(display_chunk_metadata) await processor.process_event(ReasoningStartEvent( id=self._current_reasoning_id, - metadata=chunk_metadata, + metadata=display_chunk_metadata, )) if event_type == "reasoning-end" and hasattr(self, '_current_reasoning_id'): + await _flush_reasoning_rewriter() reasoning_end_metadata = getattr(self, '_current_reasoning_metadata', {}) or {} await processor.process_event(ReasoningEndEvent( id=self._current_reasoning_id, @@ -2790,22 +2834,26 @@ def _build_llm_response_payload( if not hasattr(self, '_current_reasoning_id'): reasoning_id_counter += 1 self._current_reasoning_id = f"reasoning-{reasoning_id_counter}" - self._current_reasoning_metadata = dict(chunk_metadata) + self._current_reasoning_metadata = dict(display_chunk_metadata) await processor.process_event(ReasoningStartEvent( id=self._current_reasoning_id, - metadata=chunk_metadata, + metadata=display_chunk_metadata, )) if chunk_reasoning: - await processor.process_event(ReasoningDeltaEvent( - id=self._current_reasoning_id, - text=chunk_reasoning, - metadata=chunk_metadata, - )) + if stream_reasoning_rewriter is not None: + reasoning_text = stream_reasoning_rewriter.feed(reasoning_text) + if reasoning_text: + await processor.process_event(ReasoningDeltaEvent( + id=self._current_reasoning_id, + text=reasoning_text, + metadata=display_chunk_metadata, + )) # 2) End reasoning block when this chunk also carries non-reasoning # content (or once the stream moves away from reasoning). if (chunk_text or chunk_tool_calls) and hasattr(self, '_current_reasoning_id'): + await _flush_reasoning_rewriter() reasoning_end_metadata = getattr(self, '_current_reasoning_metadata', {}) or {} await processor.process_event(ReasoningEndEvent( id=self._current_reasoning_id, @@ -2816,8 +2864,13 @@ def _build_llm_response_payload( delattr(self, '_current_reasoning_metadata') # 3) Process text delta. - if chunk_text: + raw_chunk_text = chunk_text + if chunk_text and stream_text_rewriter is not None: + chunk_text = stream_text_rewriter.feed(chunk_text) + + if raw_chunk_text: chunk_counts["text"] += 1 + if chunk_text: if not text_started: await processor.process_event(TextStartEvent()) text_started = True @@ -2867,6 +2920,14 @@ def _build_llm_response_payload( }) await tool_accumulator.flush_remaining(stream_finish_reason) + + if stream_text_rewriter is not None: + trailing_text = stream_text_rewriter.flush() + if trailing_text: + if not text_started: + await processor.process_event(TextStartEvent()) + text_started = True + await processor.process_event(TextDeltaEvent(text=trailing_text)) # End text block if started if text_started: @@ -2874,6 +2935,7 @@ def _build_llm_response_payload( # End any remaining reasoning block if hasattr(self, '_current_reasoning_id'): + await _flush_reasoning_rewriter() reasoning_end_metadata = getattr(self, '_current_reasoning_metadata', {}) or {} await processor.process_event(ReasoningEndEvent( id=self._current_reasoning_id, diff --git a/flocks/session/streaming/stream_processor.py b/flocks/session/streaming/stream_processor.py index 1835e0ff8..6dfbd572c 100644 --- a/flocks/session/streaming/stream_processor.py +++ b/flocks/session/streaming/stream_processor.py @@ -568,6 +568,37 @@ async def _handle_tool_call(self, event: ToolCallEvent) -> None: }) return + # Hook pipeline: tool.execute.before + # Apply any input rewrite before publishing the running state so UI + # surfaces show the actual tool input that will be executed. + hook_skip = False + hook_skip_error = "Tool execution blocked by hook" + try: + from flocks.hooks.pipeline import HookPipeline + hook_ctx = await HookPipeline.run_tool_before({ + "sessionID": self.session_id, + "workspace": self._workspace_dir, + "agent": self.agent.name, + "tool": { + "name": tool_name, + "input": tool_input, + "callID": tool_call_id, + }, + }) + if hook_ctx and isinstance(hook_ctx.input, dict): + updated = hook_ctx.input.get("tool", {}).get("input") + if isinstance(updated, dict): + tool_input = updated + tool_state.input = tool_input + hook_output = hook_ctx.output if hook_ctx and isinstance(hook_ctx.output, dict) else {} + hook_skip = hook_output.get("skip", False) + if isinstance(hook_output.get("error"), str) and hook_output["error"].strip(): + hook_skip_error = hook_output["error"].strip() + except Exception as e: + log.error("stream.tool_before_hook.error", {"error": str(e)}) + hook_skip = True + hook_skip_error = "Tool execution blocked because tool-before hook failed" + tool_state.status = "running" # Update ToolPart to running state (like Flocks's Session.updatePart) @@ -656,29 +687,6 @@ async def _handle_tool_call(self, event: ToolCallEvent) -> None: await self.tool_start_callback(tool_name, tool_input) except Exception as e: log.error("stream.tool_start_callback.error", {"error": str(e)}) - - # Hook pipeline: tool.execute.before - try: - from flocks.hooks.pipeline import HookPipeline - hook_ctx = await HookPipeline.run_tool_before({ - "sessionID": self.session_id, - "workspace": self._workspace_dir, - "agent": self.agent.name, - "tool": { - "name": tool_name, - "input": tool_input, - "callID": tool_call_id, - }, - }) - if hook_ctx and isinstance(hook_ctx.input, dict): - updated = hook_ctx.input.get("tool", {}).get("input") - if isinstance(updated, dict): - tool_input = updated - hook_skip = hook_ctx.output.get("skip") if hook_ctx else False - except Exception as e: - log.error("stream.tool_before_hook.error", {"error": str(e)}) - hook_skip = False - # Execute tool synchronously tool_span_ctx = None if self._langfuse_generation is not None: @@ -702,7 +710,7 @@ async def _handle_tool_call(self, event: ToolCallEvent) -> None: if hook_skip: result = ToolResult( success=False, - error="Tool execution blocked by hook", + error=hook_skip_error, ) else: sandbox_meta = await self._resolve_sandbox_meta(tool_name) diff --git a/tests/session/test_cli_title_generation.py b/tests/session/test_cli_title_generation.py index 4693d3a8f..18643fb4d 100644 --- a/tests/session/test_cli_title_generation.py +++ b/tests/session/test_cli_title_generation.py @@ -190,6 +190,41 @@ async def failing_stream(*args, **kwargs): assert title == "Help me with something" mock_update.assert_awaited_once() + @pytest.mark.asyncio + async def test_falls_back_without_provider_call_when_llm_before_hook_fails(self): + """A failing llm_before hook blocks title provider calls and uses local fallback.""" + from flocks.session.lifecycle.title import SessionTitle + from flocks.hooks.pipeline import HookPipeline + + question = "Contact alice@example.com about the alert" + mock_session = _make_session() + msg, part = _make_user_msg(question) + mock_provider = MagicMock() + mock_provider.chat_stream = MagicMock(side_effect=AssertionError("provider must not be called")) + mock_update = AsyncMock() + + patches = _patch_title_deps(mock_session, [msg], [part], mock_provider, mock_update) + with ( + patches[0], + patches[1], + patches[2], + patches[3], + patches[4], + patches[5], + patches[6], + patch.object(HookPipeline, "has_stage_handlers", new=AsyncMock(return_value=True)), + patch.object(HookPipeline, "run_llm_before", new=AsyncMock(side_effect=RuntimeError("hook boom"))), + ): + title = await SessionTitle.generate_title_after_first_message( + session_id="sess-1", + model_id="claude-3", + provider_id="anthropic", + ) + + assert title == SessionTitle._generate_simple_title(question) + mock_provider.chat_stream.assert_not_called() + mock_update.assert_awaited_once() + @pytest.mark.asyncio async def test_publishes_sse_event_when_callback_provided(self): """SSE event is published when event_publish_callback is given (Web path).""" diff --git a/tests/session/test_compaction_iterative_summary.py b/tests/session/test_compaction_iterative_summary.py index 17c0a4117..634167219 100644 --- a/tests/session/test_compaction_iterative_summary.py +++ b/tests/session/test_compaction_iterative_summary.py @@ -25,10 +25,12 @@ ) from flocks.session.lifecycle.compaction import compaction as compaction_module from flocks.session.lifecycle.compaction.summary import ( + _llm_chat_with_timeout, build_iterative_prompt, summarize_chunked_iterative, summarize_single_pass, ) +from flocks.provider.provider import ChatMessage # --------------------------------------------------------------------------- @@ -610,3 +612,36 @@ async def test_no_previous_summary_uses_default_prompt(self) -> None: body = call.kwargs["messages"][0].content assert "<<>>" not in body assert "## Decisions" in body + + +@pytest.mark.asyncio +async def test_compaction_provider_call_blocks_when_llm_before_hook_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from flocks.hooks.pipeline import HookPipeline, HookStage + + provider = MagicMock() + provider.chat = AsyncMock(side_effect=AssertionError("provider must not be called")) + + monkeypatch.setattr( + HookPipeline, + "has_stage_handlers", + AsyncMock(side_effect=lambda stage, _metadata=None: stage == HookStage.LLM_BEFORE), + ) + monkeypatch.setattr( + HookPipeline, + "run_llm_before", + AsyncMock(side_effect=RuntimeError("hook boom")), + ) + + with pytest.raises(RuntimeError, match="request was not sent"): + await _llm_chat_with_timeout( + provider_client=provider, + model_id="test-model", + messages=[ChatMessage(role="user", content="email alice@example.com")], + max_tokens=100, + timeout=5, + session_id="ses_compaction_hook_fail", + ) + + provider.chat.assert_not_awaited() diff --git a/tests/session/test_goal.py b/tests/session/test_goal.py index 7e845d7f3..f9a538a0e 100644 --- a/tests/session/test_goal.py +++ b/tests/session/test_goal.py @@ -234,6 +234,65 @@ async def test_goal_evaluation_uses_model_judge_when_provider_model_are_availabl assert decision.reason == "The final response says the implementation and tests are complete." +@pytest.mark.asyncio +async def test_goal_model_judge_applies_llm_before_hook(monkeypatch: pytest.MonkeyPatch): + session_id = "goal_model_judge_redaction_session" + await GoalManager.set_goal(session_id, "handle alice@example.com") + provider = SimpleNamespace( + chat=AsyncMock(return_value=SimpleNamespace( + content='{"verdict": "continue", "reason": "Need more work for [[V_EMAIL_1]]."}' + )) + ) + + async def _run_llm_before(payload): + assert payload["sessionID"] == session_id + assert "alice@example.com" in payload["request"]["messages"][1]["content"] + updated_request = dict(payload["request"]) + updated_request["messages"] = [ + payload["request"]["messages"][0], + { + **payload["request"]["messages"][1], + "content": payload["request"]["messages"][1]["content"].replace( + "alice@example.com", + "[[V_EMAIL_1]]", + ), + }, + ] + return SimpleNamespace( + output={ + "request": updated_request, + "redaction": { + "streamTextReplacements": [ + {"placeholder": "[[V_EMAIL_1]]", "value": "alice@example.com"} + ] + }, + } + ) + + monkeypatch.setattr( + "flocks.hooks.pipeline.HookPipeline.has_stage_handlers", + AsyncMock(return_value=True), + ) + monkeypatch.setattr( + "flocks.hooks.pipeline.HookPipeline.run_llm_before", + AsyncMock(side_effect=_run_llm_before), + ) + + with patch("flocks.session.goal.Provider.get", return_value=provider): + decision = await GoalManager.evaluate_after_turn( + session_id, + "Latest response mentions alice@example.com.", + provider_id="test-provider", + model_id="test-model", + ) + + provider.chat.assert_awaited_once() + sent_prompt = provider.chat.await_args.kwargs["messages"][1].content + assert "alice@example.com" not in sent_prompt + assert "[[V_EMAIL_1]]" in sent_prompt + assert decision.reason == "Need more work for alice@example.com." + + @pytest.mark.asyncio async def test_goal_model_judge_receives_initial_clarification(): session_id = "goal_model_judge_clarification_session" diff --git a/tests/session/test_runner_llm_hooks.py b/tests/session/test_runner_llm_hooks.py index de691d02f..cc473ded4 100644 --- a/tests/session/test_runner_llm_hooks.py +++ b/tests/session/test_runner_llm_hooks.py @@ -59,6 +59,9 @@ def get_reasoning_content(self) -> str: def get_finish_reason(self): return self.finish_reason + async def drain_parallel_tool_calls(self) -> None: + return None + class _FakeToolAccumulator: def __init__(self, processor): @@ -163,6 +166,11 @@ async def _after(payload, result): assert result["chunkCounts"] == {"total": 1, "reasoning": 1, "text": 1, "tool": 0} monkeypatch.setattr(runner_mod, "StreamProcessor", _FakeProcessor) + monkeypatch.setattr( + runner_mod.HookPipeline, + "has_stage_handlers", + AsyncMock(return_value=True), + ) monkeypatch.setattr( runner_mod.HookPipeline, "run_llm_before", @@ -239,6 +247,174 @@ async def _gen(): assert order == ["before", "provider", "after"] +@pytest.mark.asyncio +async def test_call_llm_blocks_provider_when_llm_before_hook_fails(monkeypatch: pytest.MonkeyPatch): + runner = _make_runner("ses_runner_llm_before_fail_closed") + assistant_msg = SimpleNamespace(id="msg_assistant_before_fail") + agent = SimpleNamespace(name="rex") + + monkeypatch.setattr(runner_mod, "StreamProcessor", _FakeProcessor) + monkeypatch.setattr( + runner_mod.HookPipeline, + "has_stage_handlers", + AsyncMock(side_effect=lambda stage, _metadata=None: stage == runner_mod.HookStage.LLM_BEFORE), + ) + monkeypatch.setattr( + runner_mod.HookPipeline, + "run_llm_before", + AsyncMock(side_effect=RuntimeError("redaction unavailable")), + ) + monkeypatch.setattr( + runner_mod, + "langfuse_is_active", + lambda: False, + ) + monkeypatch.setattr( + "flocks.provider.options.build_provider_options", + lambda provider_id, model_id: {}, + ) + monkeypatch.setattr( + "flocks.session.streaming.tool_accumulator.ToolCallAccumulator", + _FakeToolAccumulator, + ) + + class _Provider: + def chat_stream(self, **kwargs): + raise AssertionError("provider must not be called when llm_before fails") + + with pytest.raises(RuntimeError, match="request was not sent"): + await runner._call_llm( + provider=_Provider(), + messages=[ChatMessage(role="user", content="email alice@example.com")], + tools=[], + agent=agent, + assistant_msg=assistant_msg, + ) + + +@pytest.mark.asyncio +async def test_call_llm_blocks_provider_when_hook_stage_probe_fails(monkeypatch: pytest.MonkeyPatch): + runner = _make_runner("ses_runner_hook_probe_fail_closed") + assistant_msg = SimpleNamespace(id="msg_assistant_probe_fail") + agent = SimpleNamespace(name="rex") + + monkeypatch.setattr(runner_mod, "StreamProcessor", _FakeProcessor) + monkeypatch.setattr( + runner_mod.HookPipeline, + "has_stage_handlers", + AsyncMock(side_effect=RuntimeError("hook registry unavailable")), + ) + monkeypatch.setattr( + runner_mod, + "langfuse_is_active", + lambda: False, + ) + monkeypatch.setattr( + "flocks.provider.options.build_provider_options", + lambda provider_id, model_id: {}, + ) + monkeypatch.setattr( + "flocks.session.streaming.tool_accumulator.ToolCallAccumulator", + _FakeToolAccumulator, + ) + + class _Provider: + def chat_stream(self, **kwargs): + raise AssertionError("provider must not be called when hook probe fails") + + with pytest.raises(RuntimeError, match="stage probe failed"): + await runner._call_llm( + provider=_Provider(), + messages=[ChatMessage(role="user", content="email alice@example.com")], + tools=[], + agent=agent, + assistant_msg=assistant_msg, + ) + + +@pytest.mark.asyncio +async def test_call_llm_initializes_langfuse_after_llm_before_redaction(monkeypatch: pytest.MonkeyPatch): + runner = _make_runner("ses_runner_langfuse_redacted") + assistant_msg = SimpleNamespace(id="msg_assistant_langfuse_redacted") + agent = SimpleNamespace(name="rex") + generation_inputs: list[list[dict[str, object]]] = [] + trace_inputs: list[dict[str, object]] = [] + + async def _before(payload): + raw_messages = payload["request"]["messages"] + assert raw_messages[0]["content"] == "email alice@example.com" + return SimpleNamespace( + output={ + "request": { + **payload["request"], + "messages": [{"role": "user", "content": "email [[V_EMAIL_1]]"}], + "providerOptions": {}, + }, + "redaction": { + "streamTextReplacements": [ + {"placeholder": "[[V_EMAIL_1]]", "value": "alice@example.com"} + ], + }, + } + ) + + def _trace_scope(**kwargs): + trace_inputs.append(kwargs["input"]) + return SimpleNamespace(observation="trace") + + def _generation_scope(**kwargs): + generation_inputs.append(kwargs["input"]) + return SimpleNamespace(observation="generation") + + monkeypatch.setattr(runner_mod, "StreamProcessor", _FakeProcessor) + monkeypatch.setattr( + runner_mod.HookPipeline, + "has_stage_handlers", + AsyncMock(side_effect=lambda stage, _metadata=None: stage == runner_mod.HookStage.LLM_BEFORE), + ) + monkeypatch.setattr( + runner_mod.HookPipeline, + "run_llm_before", + AsyncMock(side_effect=_before), + ) + monkeypatch.setattr(runner_mod, "langfuse_is_active", lambda: True) + monkeypatch.setattr(runner_mod, "trace_scope", _trace_scope) + monkeypatch.setattr(runner_mod, "generation_scope", _generation_scope) + monkeypatch.setattr( + "flocks.provider.options.build_provider_options", + lambda provider_id, model_id: {}, + ) + monkeypatch.setattr( + "flocks.session.streaming.tool_accumulator.ToolCallAccumulator", + _FakeToolAccumulator, + ) + monkeypatch.setattr(runner_mod.Message, "update", AsyncMock(return_value=None)) + + class _Provider: + def chat_stream(self, **kwargs): + assert kwargs["messages"][0].content == "email [[V_EMAIL_1]]" + + async def _gen(): + yield SimpleNamespace(delta="done", finish_reason="stop") + + return _gen() + + result = await runner._call_llm( + provider=_Provider(), + messages=[ChatMessage(role="user", content="email alice@example.com")], + tools=[], + agent=agent, + assistant_msg=assistant_msg, + ) + + assert result.action == "stop" + assert generation_inputs + assert trace_inputs + assert "alice@example.com" not in str(generation_inputs) + assert "alice@example.com" not in str(trace_inputs) + assert "[[V_EMAIL_1]]" in str(generation_inputs) + + @pytest.mark.asyncio async def test_call_llm_emits_after_hook_on_error(monkeypatch: pytest.MonkeyPatch): runner = _make_runner("ses_runner_llm_hooks_error") @@ -258,6 +434,11 @@ async def _after(payload, result): assert "provider boom" in result["error"]["message"] monkeypatch.setattr(runner_mod, "StreamProcessor", _FakeProcessor) + monkeypatch.setattr( + runner_mod.HookPipeline, + "has_stage_handlers", + AsyncMock(return_value=True), + ) monkeypatch.setattr( runner_mod.HookPipeline, "run_llm_before", diff --git a/tests/session/test_stream_processor.py b/tests/session/test_stream_processor.py index 186de7c3f..00818b308 100644 --- a/tests/session/test_stream_processor.py +++ b/tests/session/test_stream_processor.py @@ -449,6 +449,84 @@ async def _fake_execute(*, tool_name, ctx, **kwargs): assert seen_abort["event"] is abort_event + @pytest.mark.asyncio + async def test_tool_before_rewrite_is_published_in_running_state(self): + event_callback = AsyncMock() + proc = _make_processor(event_callback=event_callback) + + async def _fake_tool_before(payload): + assert payload["tool"]["input"] == {"ip": "[IP_1]"} + payload["tool"]["input"] = {"ip": "10.1.2.3"} + ctx = MagicMock() + ctx.input = payload + ctx.output = {} + return ctx + + result = ToolResult(success=True, output="ok", title="ip query", metadata={}) + + with ( + patch("flocks.session.streaming.stream_processor.Message.store_part", new=AsyncMock()), + patch("flocks.session.streaming.stream_processor.Message.update_part", new=AsyncMock()), + patch( + "flocks.session.streaming.stream_processor.ToolRegistry.execute", + new=AsyncMock(return_value=result), + ), + patch( + "flocks.hooks.pipeline.HookPipeline.run_tool_before", + new=AsyncMock(side_effect=_fake_tool_before), + ), + ): + await proc.process_event(ToolInputStartEvent(id="tc_restore_running", tool_name="ip_query")) + await proc.process_event( + ToolCallEvent( + tool_call_id="tc_restore_running", + tool_name="ip_query", + input={"ip": "[IP_1]"}, + ) + ) + + running_inputs = [ + call.args[1]["part"]["state"]["input"] + for call in event_callback.await_args_list + if ( + call.args[0] == "message.part.updated" + and call.args[1].get("part", {}).get("type") == "tool" + and call.args[1]["part"]["state"]["status"] == "running" + ) + ] + assert running_inputs == [{"ip": "10.1.2.3"}] + + @pytest.mark.asyncio + async def test_tool_before_failure_blocks_tool_execution(self): + proc = _make_processor() + execute_mock = AsyncMock(return_value=ToolResult(success=True, output="should not run", title="ip query")) + + with ( + patch("flocks.session.streaming.stream_processor.Message.store_part", new=AsyncMock()), + patch("flocks.session.streaming.stream_processor.Message.update_part", new=AsyncMock()), + patch( + "flocks.session.streaming.stream_processor.ToolRegistry.execute", + new=execute_mock, + ), + patch( + "flocks.hooks.pipeline.HookPipeline.run_tool_before", + new=AsyncMock(side_effect=RuntimeError("redaction restore failed")), + ), + ): + await proc.process_event(ToolInputStartEvent(id="tc_hook_fail", tool_name="ip_query")) + await proc.process_event( + ToolCallEvent( + tool_call_id="tc_hook_fail", + tool_name="ip_query", + input={"ip": "[[V_IP_ADDRESS_1]]"}, + ) + ) + + execute_mock.assert_not_awaited() + state = proc.tool_calls["tc_hook_fail"] + assert state.status == "error" + assert state.error == "Tool execution blocked because tool-before hook failed" + @pytest.mark.asyncio async def test_cancelled_tool_blocks_late_running_metadata_updates(self): event_callback = AsyncMock() diff --git a/webui/src/api/sensitiveDetection.ts b/webui/src/api/sensitiveDetection.ts new file mode 100644 index 000000000..cc4f23c79 --- /dev/null +++ b/webui/src/api/sensitiveDetection.ts @@ -0,0 +1,41 @@ +import client from './client'; + +export type PromptRedactionPlaceholderFormat = 'verbose' | 'compact'; + +export interface PromptRedactionSettings { + enabled: boolean; + categories: string[] | null; + placeholderFormat: PromptRedactionPlaceholderFormat; + promptHintEnabled: boolean; +} + +export interface PromptRedactionSettingsUpdate { + enabled?: boolean; + categories?: string[] | null; + placeholderFormat?: PromptRedactionPlaceholderFormat; + promptHintEnabled?: boolean; +} + +function normalizeSettings(data: any): PromptRedactionSettings { + const raw = data?.data ?? data ?? {}; + return { + enabled: raw.enabled === true, + categories: Array.isArray(raw.categories) ? raw.categories : null, + placeholderFormat: raw.placeholderFormat === 'compact' ? 'compact' : 'verbose', + promptHintEnabled: raw.promptHintEnabled !== false, + }; +} + +export const sensitiveDetectionApi = { + getPromptRedactionSettings: async (): Promise => { + const response = await client.get('/api/flockspro/sensitive-detection/settings'); + return normalizeSettings(response.data); + }, + + updatePromptRedactionSettings: async ( + payload: PromptRedactionSettingsUpdate, + ): Promise => { + const response = await client.patch('/api/flockspro/sensitive-detection/settings', payload); + return normalizeSettings(response.data); + }, +}; diff --git a/webui/src/locales/en-US/flockspro.json b/webui/src/locales/en-US/flockspro.json index 08bce3525..9b0735ec1 100644 --- a/webui/src/locales/en-US/flockspro.json +++ b/webui/src/locales/en-US/flockspro.json @@ -154,6 +154,28 @@ "invalidEmailError": "Please enter a valid applicant email.", "invalidPhoneError": "Please enter a valid applicant phone number (international numbers supported)." }, + "sensitiveDetection": { + "title": "Prompt Redaction", + "description": "When enabled, sensitive values in user, assistant, and tool messages are replaced before model input.", + "enabledLabel": "Enable prompt redaction", + "placeholderFormat": "Placeholder Format", + "verbosePreview": "Example: [[V_EMAIL_1]]", + "compactPreview": "Example: [EMAIL_1]", + "promptHint": "Model Value Passing Hint", + "promptHintDescription": "Add guidance for passing replaced values to tools. Recommended on; disabling it may make models treat values as template variables.", + "loading": "Loading prompt redaction settings...", + "saving": "Saving prompt redaction settings...", + "placeholderFormats": { + "verbose": "Stable", + "compact": "Compact" + }, + "errors": { + "fetch": "Failed to load prompt redaction settings", + "save": "Failed to save prompt redaction settings", + "fetchSettings": "Failed to load sensitive detection settings", + "updateSettings": "Failed to save sensitive detection settings" + } + }, "callback": { "processing": "Completing console login, please wait...", "missingConsoleLoginId": "Missing console_login_id in callback URL.", diff --git a/webui/src/locales/zh-CN/flockspro.json b/webui/src/locales/zh-CN/flockspro.json index d0d9829b1..53b2c239e 100644 --- a/webui/src/locales/zh-CN/flockspro.json +++ b/webui/src/locales/zh-CN/flockspro.json @@ -154,6 +154,28 @@ "invalidEmailError": "请输入有效的申请人邮箱", "invalidPhoneError": "请输入有效的申请人电话(支持国际号码)" }, + "sensitiveDetection": { + "title": "Prompt 脱敏", + "description": "开启后,进入大模型前会替换用户、助手和工具消息中的敏感信息。", + "enabledLabel": "启用 Prompt 脱敏", + "placeholderFormat": "占位符格式", + "verbosePreview": "示例:[[V_EMAIL_1]]", + "compactPreview": "示例:[EMAIL_1]", + "promptHint": "模型值传递提示", + "promptHintDescription": "向模型补充如何把替换值传递给工具。建议开启,关闭后模型可能把替换值误判为模板变量。", + "loading": "正在读取 Prompt 脱敏设置...", + "saving": "正在保存 Prompt 脱敏设置...", + "placeholderFormats": { + "verbose": "稳定格式", + "compact": "简短格式" + }, + "errors": { + "fetch": "读取 Prompt 脱敏设置失败", + "save": "保存 Prompt 脱敏设置失败", + "fetchSettings": "读取敏感信息识别设置失败", + "updateSettings": "保存敏感信息识别设置失败" + } + }, "callback": { "processing": "正在完成云账号登录,请稍候...", "missingConsoleLoginId": "缺少 console_login_id,无法完成登录。", diff --git a/webui/src/pages/FlocksproUpgrade/index.tsx b/webui/src/pages/FlocksproUpgrade/index.tsx index 89eb4198f..2c957626d 100644 --- a/webui/src/pages/FlocksproUpgrade/index.tsx +++ b/webui/src/pages/FlocksproUpgrade/index.tsx @@ -1,5 +1,5 @@ import { useCallback, useEffect, useMemo, useRef, useState } from 'react'; -import { ArrowDownCircle, ArrowUpCircle, CheckCircle, ChevronDown, Loader2, LogIn, X, XCircle } from 'lucide-react'; +import { ArrowDownCircle, ArrowUpCircle, CheckCircle, ChevronDown, Loader2, LogIn, ShieldCheck, X, XCircle } from 'lucide-react'; import { useSearchParams } from 'react-router-dom'; import { useTranslation } from 'react-i18next'; import PageHeader from '@/components/common/PageHeader'; @@ -12,6 +12,12 @@ import { type UpgradeRequestStatus, } from '@/api/consoleUpgrade'; import { useProductName } from '@/contexts/ProductNameContext'; +import { + sensitiveDetectionApi, + type PromptRedactionPlaceholderFormat, + type PromptRedactionSettings, + type PromptRedactionSettingsUpdate, +} from '@/api/sensitiveDetection'; import { type UpdateProgress } from '@/api/update'; import { extractErrorMessage } from '@/utils/error'; import { checkRestartReadiness } from '@/utils/restartPolling'; @@ -329,6 +335,10 @@ export default function FlocksproUpgradePage() { const [showLicenseDetails, setShowLicenseDetails] = useState(false); const [licenseStatus, setLicenseStatus] = useState(null); const [proPackageStatus, setProPackageStatus] = useState(null); + const [promptRedactionSettings, setPromptRedactionSettings] = useState(null); + const [promptRedactionLoading, setPromptRedactionLoading] = useState(false); + const [promptRedactionSaving, setPromptRedactionSaving] = useState(false); + const [promptRedactionError, setPromptRedactionError] = useState(null); const [dismissedRejectedRequestIds, setDismissedRejectedRequestIds] = useState>( loadDismissedRejectedRequestIds, ); @@ -792,6 +802,47 @@ export default function FlocksproUpgradePage() { } }, [activeRequest?.request_id, currentDisplayLicenseRequest?.request_id, currentIssuedRequest?.request_id, refreshRequests, t]); + const loadPromptRedactionSettings = useCallback(async () => { + if (!isProLoaded) { + setPromptRedactionSettings(null); + setPromptRedactionError(null); + return; + } + setPromptRedactionLoading(true); + setPromptRedactionError(null); + try { + const settings = await sensitiveDetectionApi.getPromptRedactionSettings(); + setPromptRedactionSettings(settings); + } catch (err) { + setPromptRedactionError(extractErrorMessage(err, t('sensitiveDetection.errors.fetchSettings'))); + } finally { + setPromptRedactionLoading(false); + } + }, [isProLoaded, t]); + + const updatePromptRedactionSettings = useCallback( + async (payload: PromptRedactionSettingsUpdate) => { + if (!isProLoaded) { + return; + } + setPromptRedactionSaving(true); + setPromptRedactionError(null); + try { + const settings = await sensitiveDetectionApi.updatePromptRedactionSettings(payload); + setPromptRedactionSettings(settings); + } catch (err) { + setPromptRedactionError(extractErrorMessage(err, t('sensitiveDetection.errors.updateSettings'))); + } finally { + setPromptRedactionSaving(false); + } + }, + [isProLoaded, t], + ); + + useEffect(() => { + void loadPromptRedactionSettings(); + }, [loadPromptRedactionSettings]); + useEffect(() => { if (autoSyncTriggeredRef.current) { return; @@ -1319,6 +1370,112 @@ export default function FlocksproUpgradePage() { )} + {isProLoaded && ( +
+
+
+ + + +
+

{t('sensitiveDetection.title')}

+

{t('sensitiveDetection.description')}

+
+
+ +
+ +
+
+
+
{t('sensitiveDetection.placeholderFormat')}
+
+ {promptRedactionSettings?.placeholderFormat === 'compact' + ? t('sensitiveDetection.compactPreview') + : t('sensitiveDetection.verbosePreview')} +
+
+
+ {(['verbose', 'compact'] as PromptRedactionPlaceholderFormat[]).map((format) => { + const active = promptRedactionSettings?.placeholderFormat === format; + return ( + + ); + })} +
+
+ +
+
+
{t('sensitiveDetection.promptHint')}
+
{t('sensitiveDetection.promptHintDescription')}
+
+ +
+
+ + {(promptRedactionLoading || promptRedactionSaving || promptRedactionError) && ( +
+ {promptRedactionError || + (promptRedactionSaving ? t('sensitiveDetection.saving') : t('sensitiveDetection.loading'))} +
+ )} +
+ )} + {historyRequests.length > 0 && (