tokio_prompt_orchestrator/plugin.rs
1//! Custom plugin stage system for the LLM inference pipeline.
2//!
3//! This module provides a first-class extension point so user code can inject
4//! arbitrary async processing logic into the pipeline without forking the
5//! library. Plugins are composable, ordered, and fully type-safe.
6//!
7//! ## Architecture
8//!
9//! ```text
10//! PromptRequest
11//! |
12//! v
13//! [BeforeRag plugins] -> RAG stage -> [AfterRag plugins]
14//! |
15//! v
16//! [BeforeAssemble plugins] -> Assemble stage -> [AfterAssemble plugins]
17//! |
18//! v
19//! [BeforeInference plugins] -> Inference stage -> [AfterInference plugins]
20//! |
21//! v
22//! [BeforePost plugins] -> Post stage -> [AfterPost plugins]
23//! |
24//! v
25//! [BeforeStream plugins] -> Stream stage -> [AfterStream plugins]
26//! ```
27//!
28//! ## Quick Start
29//!
30//! ```rust,no_run
31//! use tokio_prompt_orchestrator::plugin::{
32//! StagePlugin, PluginInput, PluginOutput, PluginPosition, PluginRegistry,
33//! };
34//! use tokio_prompt_orchestrator::PipelineStage;
35//! use async_trait::async_trait;
36//! use std::sync::Arc;
37//!
38//! /// A simple logging plugin that records when the inference stage is entered.
39//! struct InferenceLogPlugin;
40//!
41//! #[async_trait]
42//! impl StagePlugin for InferenceLogPlugin {
43//! fn name(&self) -> &'static str { "inference-logger" }
44//!
45//! async fn process(&self, input: PluginInput) -> PluginOutput {
46//! tracing::info!(request_id = %input.request_id, "entering inference stage");
47//! PluginOutput::passthrough(input)
48//! }
49//! }
50//!
51//! #[tokio::main]
52//! async fn main() {
53//! let mut registry = PluginRegistry::new();
54//! registry.register(
55//! PluginPosition::Before(PipelineStage::Inference),
56//! Arc::new(InferenceLogPlugin),
57//! );
58//!
59//! let chain = registry.chain_for(PluginPosition::Before(PipelineStage::Inference));
60//! let input = PluginInput {
61//! request_id: "req-1".to_string(),
62//! session_id: "session-1".to_string(),
63//! payload: serde_json::json!({"prompt": "hello"}),
64//! metadata: std::collections::HashMap::new(),
65//! };
66//!
67//! let output = chain.run(input).await;
68//! println!("Plugin chain produced: {:?}", output.status);
69//! }
70//! ```
71
72use crate::PipelineStage;
73use async_trait::async_trait;
74use serde_json::Value;
75use std::collections::HashMap;
76use std::sync::Arc;
77use std::time::Duration;
78use tracing::{debug, warn};
79
80// ============================================================================
81// PluginInput / PluginOutput
82// ============================================================================
83
84/// Input passed to every plugin in a [`PluginChain`].
85///
86/// The `payload` field carries stage-specific data as an untyped JSON value so
87/// that the same trait works across all pipeline stages without requiring a
88/// separate trait per stage.
89///
90/// # Cloning
91///
92/// `PluginInput` is cheap to clone; `payload` clones are backed by `Arc`
93/// internally in `serde_json::Value`.
94#[derive(Debug, Clone)]
95pub struct PluginInput {
96 /// Unique request ID for distributed trace correlation.
97 pub request_id: String,
98 /// Session this request belongs to.
99 pub session_id: String,
100 /// Stage-specific payload (prompt text, tokens, assembled prompt, …).
101 ///
102 /// Plugins that do not need to inspect or mutate the payload should pass
103 /// it through unchanged via [`PluginOutput::passthrough`].
104 pub payload: Value,
105 /// Arbitrary key-value metadata forwarded from the originating
106 /// [`crate::PromptRequest`]. Plugins may read and extend this map.
107 pub metadata: HashMap<String, String>,
108}
109
110/// The outcome status of a plugin's `process` call.
111///
112/// Consumers (typically [`PluginChain::run`]) inspect this to decide whether
113/// to continue executing the remaining plugins in the chain.
114#[derive(Debug, Clone, PartialEq, Eq)]
115pub enum PluginStatus {
116 /// Processing may continue normally.
117 Continue,
118 /// This plugin wants to short-circuit the remaining chain.
119 ///
120 /// `PluginChain::run` stops executing subsequent plugins and returns this
121 /// output immediately. The pipeline stage itself still runs; use this to
122 /// replace or skip preprocessing steps, not to cancel inference entirely.
123 Abort,
124 /// The plugin detected an error and wants to propagate it.
125 ///
126 /// Carries a human-readable description. Like `Abort`, the remaining
127 /// chain is skipped. The caller receives the `Error` status and can
128 /// choose to shed the request to the dead-letter queue.
129 Error(String),
130}
131
132/// Output returned by a [`StagePlugin::process`] call.
133#[derive(Debug, Clone)]
134pub struct PluginOutput {
135 /// The (possibly modified) input to pass to the next plugin or stage.
136 pub input: PluginInput,
137 /// Status signal for the chain runner.
138 pub status: PluginStatus,
139}
140
141impl PluginOutput {
142 /// Construct a `Continue` output that passes `input` through unchanged.
143 ///
144 /// This is the most common return value for plugins that only observe
145 /// or annotate requests without modifying the payload.
146 ///
147 /// # Panics
148 ///
149 /// This function does not panic.
150 #[must_use]
151 pub fn passthrough(input: PluginInput) -> Self {
152 Self {
153 input,
154 status: PluginStatus::Continue,
155 }
156 }
157
158 /// Construct a `Continue` output with a mutated payload.
159 ///
160 /// Use this when a plugin needs to transform the data before it reaches
161 /// the next stage (e.g., sanitising PII, injecting system context).
162 ///
163 /// # Panics
164 ///
165 /// This function does not panic.
166 #[must_use]
167 pub fn modified(mut input: PluginInput, payload: Value) -> Self {
168 input.payload = payload;
169 Self {
170 input,
171 status: PluginStatus::Continue,
172 }
173 }
174
175 /// Construct an `Abort` output, stopping the plugin chain.
176 ///
177 /// # Panics
178 ///
179 /// This function does not panic.
180 #[must_use]
181 pub fn abort(input: PluginInput) -> Self {
182 Self {
183 input,
184 status: PluginStatus::Abort,
185 }
186 }
187
188 /// Construct an `Error` output with a descriptive message.
189 ///
190 /// # Panics
191 ///
192 /// This function does not panic.
193 #[must_use]
194 pub fn error(input: PluginInput, message: impl Into<String>) -> Self {
195 Self {
196 input,
197 status: PluginStatus::Error(message.into()),
198 }
199 }
200}
201
202// ============================================================================
203// StagePlugin trait
204// ============================================================================
205
206/// Core trait for custom pipeline stage plugins.
207///
208/// Implementors are inserted at specific points in the pipeline via a
209/// [`PluginRegistry`]. Each plugin receives a [`PluginInput`] and must return
210/// a [`PluginOutput`] indicating whether processing should continue.
211///
212/// ## Correctness Requirements
213///
214/// - Implementations **must not** panic — return [`PluginOutput::error`]
215/// instead.
216/// - Implementations **must** be `Send + Sync` for safe use across Tokio tasks.
217/// - The `process` method **must** complete in bounded time; use
218/// [`tokio::time::timeout`] internally if calling external services.
219///
220/// ## Example
221///
222/// ```rust
223/// use tokio_prompt_orchestrator::plugin::{StagePlugin, PluginInput, PluginOutput};
224/// use async_trait::async_trait;
225///
226/// /// Redacts email addresses from the prompt payload.
227/// struct PiiRedactor;
228///
229/// #[async_trait]
230/// impl StagePlugin for PiiRedactor {
231/// fn name(&self) -> &'static str { "pii-redactor" }
232///
233/// async fn process(&self, input: PluginInput) -> PluginOutput {
234/// // Inspect and potentially mutate the payload here.
235/// PluginOutput::passthrough(input)
236/// }
237/// }
238/// ```
239#[async_trait]
240pub trait StagePlugin: Send + Sync {
241 /// A unique, human-readable name used in logs and metrics.
242 ///
243 /// Should be `kebab-case`, e.g. `"pii-redactor"` or `"context-injector"`.
244 fn name(&self) -> &'static str;
245
246 /// Process `input` and return the (potentially modified) output.
247 ///
248 /// # Errors
249 ///
250 /// Return [`PluginOutput::error`] rather than propagating panics or
251 /// returning unexpected results.
252 ///
253 /// # Panics
254 ///
255 /// Implementations must never panic.
256 async fn process(&self, input: PluginInput) -> PluginOutput;
257}
258
259// ============================================================================
260// PluginPosition
261// ============================================================================
262
263/// Specifies where in the pipeline a plugin should be inserted.
264///
265/// Each variant corresponds to a point immediately before or after one of the
266/// five built-in pipeline stages.
267#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
268pub enum PluginPosition {
269 /// Run the plugin before the named stage.
270 Before(PipelineStage),
271 /// Run the plugin after the named stage.
272 After(PipelineStage),
273}
274
275impl PluginPosition {
276 /// Return the stage this position is relative to.
277 ///
278 /// # Panics
279 ///
280 /// This function does not panic.
281 pub fn stage(&self) -> PipelineStage {
282 match self {
283 Self::Before(s) | Self::After(s) => *s,
284 }
285 }
286
287 /// Return `true` if this position is before a stage.
288 ///
289 /// # Panics
290 ///
291 /// This function does not panic.
292 pub fn is_before(&self) -> bool {
293 matches!(self, Self::Before(_))
294 }
295
296 /// Return `true` if this position is after a stage.
297 ///
298 /// # Panics
299 ///
300 /// This function does not panic.
301 pub fn is_after(&self) -> bool {
302 matches!(self, Self::After(_))
303 }
304
305 /// Return a human-readable label, e.g. `"before:inference"`.
306 ///
307 /// # Panics
308 ///
309 /// This function does not panic.
310 pub fn label(&self) -> String {
311 let prefix = if self.is_before() { "before" } else { "after" };
312 format!("{}:{}", prefix, self.stage())
313 }
314}
315
316// ============================================================================
317// PluginChain
318// ============================================================================
319
320/// An ordered sequence of [`StagePlugin`]s that run at a single
321/// [`PluginPosition`].
322///
323/// Created by [`PluginRegistry::chain_for`] and executed by calling
324/// [`PluginChain::run`]. An empty chain is a no-op: `run` returns
325/// `PluginOutput::passthrough(input)` immediately.
326///
327/// ## Execution Model
328///
329/// Plugins run serially in insertion order. If any plugin returns a
330/// non-`Continue` status, execution halts and that output is returned
331/// immediately without running the remaining plugins.
332///
333/// This matches the common middleware short-circuit pattern and keeps
334/// plugin authors free from reasoning about concurrent modifications.
335#[derive(Clone, Default)]
336pub struct PluginChain {
337 plugins: Vec<Arc<dyn StagePlugin>>,
338 position: Option<PluginPosition>,
339}
340
341impl PluginChain {
342 /// Create a new empty chain.
343 ///
344 /// # Panics
345 ///
346 /// This function does not panic.
347 #[must_use]
348 pub fn new(position: PluginPosition) -> Self {
349 Self {
350 plugins: Vec::new(),
351 position: Some(position),
352 }
353 }
354
355 /// Append a plugin to the end of the chain.
356 ///
357 /// Plugins execute in the order they were added.
358 ///
359 /// # Panics
360 ///
361 /// This function does not panic.
362 pub fn push(&mut self, plugin: Arc<dyn StagePlugin>) {
363 debug!(
364 plugin = plugin.name(),
365 position = ?self.position,
366 "registered plugin in chain"
367 );
368 self.plugins.push(plugin);
369 }
370
371 /// Return the number of plugins currently in the chain.
372 ///
373 /// # Panics
374 ///
375 /// This function does not panic.
376 pub fn len(&self) -> usize {
377 self.plugins.len()
378 }
379
380 /// Return `true` if no plugins are registered in this chain.
381 ///
382 /// # Panics
383 ///
384 /// This function does not panic.
385 pub fn is_empty(&self) -> bool {
386 self.plugins.is_empty()
387 }
388
389 /// Execute all plugins in order, stopping early on `Abort` or `Error`.
390 ///
391 /// Returns the final [`PluginOutput`] after all plugins have run (or the
392 /// first non-`Continue` output). If the chain is empty, returns
393 /// `PluginOutput::passthrough(input)` immediately.
394 ///
395 /// # Panics
396 ///
397 /// This function does not panic (plugin panics are the caller's bug).
398 pub async fn run(&self, mut input: PluginInput) -> PluginOutput {
399 if self.plugins.is_empty() {
400 return PluginOutput::passthrough(input);
401 }
402
403 let position_label = self
404 .position
405 .map(|p| p.label())
406 .unwrap_or_else(|| "unknown".to_string());
407
408 for plugin in &self.plugins {
409 let name = plugin.name();
410 debug!(plugin = name, position = %position_label, "running plugin");
411
412 let output = plugin.process(input).await;
413
414 match &output.status {
415 PluginStatus::Continue => {
416 // Propagate (possibly mutated) input to the next plugin.
417 input = output.input;
418 }
419 PluginStatus::Abort => {
420 warn!(
421 plugin = name,
422 position = %position_label,
423 "plugin aborted chain"
424 );
425 return output;
426 }
427 PluginStatus::Error(msg) => {
428 warn!(
429 plugin = name,
430 position = %position_label,
431 error = %msg,
432 "plugin returned error, aborting chain"
433 );
434 return output;
435 }
436 }
437 }
438
439 PluginOutput::passthrough(input)
440 }
441}
442
443// ============================================================================
444// PluginContext — passed to pre/post inference hooks
445// ============================================================================
446
447/// Metrics snapshot passed to pre- and post-inference hooks.
448///
449/// Values are best-effort estimates at the time the hook fires. Hooks must
450/// not modify any field — the struct is `Clone` for convenience only.
451#[derive(Debug, Clone)]
452pub struct HookMetrics {
453 /// Elapsed wall-clock time since the request entered the pipeline (ms).
454 pub elapsed_ms: u64,
455 /// Approximate number of prompt tokens (provider-reported or estimated).
456 pub prompt_tokens: Option<u64>,
457 /// Approximate number of completion tokens.
458 pub completion_tokens: Option<u64>,
459 /// Estimated cost of this request in USD (0.0 if not available).
460 pub estimated_cost_usd: f64,
461}
462
463/// Rich context passed to every pre- and post-inference hook.
464///
465/// Hooks receive a shared reference to this struct; they cannot modify the
466/// pipeline payload directly — they should use the returned [`PluginOutput`]
467/// mechanism for that. The hook context is purely informational.
468#[derive(Debug, Clone)]
469pub struct PluginContext {
470 /// Unique ID for the in-flight request.
471 pub request_id: String,
472 /// Session the request belongs to.
473 pub session_id: String,
474 /// The raw prompt/payload as submitted to the pipeline stage.
475 pub request_payload: Value,
476 /// The model response, if available (only set for post-inference hooks).
477 pub response_payload: Option<Value>,
478 /// Performance / cost metrics for this request.
479 pub metrics: HookMetrics,
480 /// Arbitrary metadata forwarded from the original [`crate::PromptRequest`].
481 pub metadata: HashMap<String, String>,
482}
483
484impl PluginContext {
485 /// Build a pre-inference context (no response yet).
486 ///
487 /// # Panics
488 ///
489 /// This function does not panic.
490 #[must_use]
491 pub fn pre_inference(
492 request_id: impl Into<String>,
493 session_id: impl Into<String>,
494 request_payload: Value,
495 metadata: HashMap<String, String>,
496 elapsed_ms: u64,
497 ) -> Self {
498 Self {
499 request_id: request_id.into(),
500 session_id: session_id.into(),
501 request_payload,
502 response_payload: None,
503 metrics: HookMetrics {
504 elapsed_ms,
505 prompt_tokens: None,
506 completion_tokens: None,
507 estimated_cost_usd: 0.0,
508 },
509 metadata,
510 }
511 }
512
513 /// Build a post-inference context (response is available).
514 ///
515 /// # Panics
516 ///
517 /// This function does not panic.
518 #[must_use]
519 #[allow(clippy::too_many_arguments)]
520 pub fn post_inference(
521 request_id: impl Into<String>,
522 session_id: impl Into<String>,
523 request_payload: Value,
524 response_payload: Value,
525 metadata: HashMap<String, String>,
526 elapsed_ms: u64,
527 prompt_tokens: Option<u64>,
528 completion_tokens: Option<u64>,
529 estimated_cost_usd: f64,
530 ) -> Self {
531 Self {
532 request_id: request_id.into(),
533 session_id: session_id.into(),
534 request_payload,
535 response_payload: Some(response_payload),
536 metrics: HookMetrics {
537 elapsed_ms,
538 prompt_tokens,
539 completion_tokens,
540 estimated_cost_usd,
541 },
542 metadata,
543 }
544 }
545
546 /// Convert this context into a [`PluginInput`] for use with `PluginChain::run`.
547 ///
548 /// The `request_payload` becomes the chain `payload`; `response_payload`
549 /// (if present) is injected into `metadata` under the key `"response_payload_json"`.
550 ///
551 /// # Panics
552 ///
553 /// This function does not panic.
554 #[must_use]
555 pub fn into_plugin_input(self) -> PluginInput {
556 let mut metadata = self.metadata;
557 if let Some(resp) = self.response_payload {
558 metadata.insert(
559 "response_payload_json".to_string(),
560 resp.to_string(),
561 );
562 }
563 metadata.insert(
564 "elapsed_ms".to_string(),
565 self.metrics.elapsed_ms.to_string(),
566 );
567 if let Some(pt) = self.metrics.prompt_tokens {
568 metadata.insert("prompt_tokens".to_string(), pt.to_string());
569 }
570 if let Some(ct) = self.metrics.completion_tokens {
571 metadata.insert("completion_tokens".to_string(), ct.to_string());
572 }
573 PluginInput {
574 request_id: self.request_id,
575 session_id: self.session_id,
576 payload: self.request_payload,
577 metadata,
578 }
579 }
580}
581
582// ============================================================================
583// InferenceHook trait — simpler API for pre/post hooks
584// ============================================================================
585
586/// A focused hook trait for code that only needs to intercept inference
587/// before or after it happens, without needing the full [`StagePlugin`]
588/// position system.
589///
590/// Hooks are registered via [`PluginRegistry::register_pre_hook`] and
591/// [`PluginRegistry::register_post_hook`]. They are called with a rich
592/// [`PluginContext`] rather than the generic [`PluginInput`].
593#[async_trait]
594pub trait InferenceHook: Send + Sync {
595 /// Human-readable name for logging/metrics.
596 fn name(&self) -> &'static str;
597
598 /// Called with the full inference context. Return `Ok(())` to continue;
599 /// return `Err(msg)` to abort the pipeline and return the error to the caller.
600 ///
601 /// # Errors
602 ///
603 /// Return an `Err` string to signal that the hook detected an unrecoverable
604 /// problem (e.g., budget exceeded, PII detected).
605 async fn call(&self, ctx: &PluginContext) -> Result<(), String>;
606}
607
608// ============================================================================
609// PluginRegistry — extended with pre/post hook support
610// ============================================================================
611
612/// Global runtime registry for [`StagePlugin`] instances.
613///
614/// Stores plugins keyed by [`PluginPosition`] and vends fully-built
615/// [`PluginChain`]s on demand. The registry is designed to be populated once
616/// at startup and then used read-only during request processing.
617///
618/// For dynamic plugin management (add/remove while the pipeline is running),
619/// wrap the registry in an `Arc<tokio::sync::RwLock<PluginRegistry>>`.
620///
621/// ## Example
622///
623/// ```rust,no_run
624/// use tokio_prompt_orchestrator::plugin::{
625/// PluginRegistry, PluginPosition, StagePlugin, PluginInput, PluginOutput,
626/// };
627/// use tokio_prompt_orchestrator::PipelineStage;
628/// use async_trait::async_trait;
629/// use std::sync::Arc;
630///
631/// struct Noop;
632///
633/// #[async_trait]
634/// impl StagePlugin for Noop {
635/// fn name(&self) -> &'static str { "noop" }
636/// async fn process(&self, input: PluginInput) -> PluginOutput {
637/// PluginOutput::passthrough(input)
638/// }
639/// }
640///
641/// let mut registry = PluginRegistry::new();
642/// registry.register(
643/// PluginPosition::Before(PipelineStage::Rag),
644/// Arc::new(Noop),
645/// );
646///
647/// assert_eq!(registry.plugin_count(PluginPosition::Before(PipelineStage::Rag)), 1);
648/// registry.remove_all(PluginPosition::Before(PipelineStage::Rag));
649/// assert_eq!(registry.plugin_count(PluginPosition::Before(PipelineStage::Rag)), 0);
650/// ```
651#[derive(Default)]
652pub struct PluginRegistry {
653 chains: HashMap<PluginPositionKey, PluginChain>,
654 /// Pre-inference hooks, executed in registration order before inference.
655 pre_hooks: Vec<(String, Arc<dyn InferenceHook>)>,
656 /// Post-inference hooks, executed in registration order after inference.
657 post_hooks: Vec<(String, Arc<dyn InferenceHook>)>,
658}
659
660/// Stable key type used as the `HashMap` key for [`PluginPosition`].
661///
662/// `PluginPosition` cannot be used directly as a map key because `PartialEq`
663/// for the enum variant alone is insufficient when hashing; this wrapper
664/// provides `Eq + Hash` using the variant's discriminant.
665#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
666struct PluginPositionKey(PluginPosition);
667
668impl PluginRegistry {
669 /// Create a new, empty registry.
670 ///
671 /// # Panics
672 ///
673 /// This function does not panic.
674 #[must_use]
675 pub fn new() -> Self {
676 Self::default()
677 }
678
679 /// Register a plugin at the given position.
680 ///
681 /// Multiple plugins registered at the same position are chained in
682 /// insertion order.
683 ///
684 /// # Panics
685 ///
686 /// This function does not panic.
687 pub fn register(&mut self, position: PluginPosition, plugin: Arc<dyn StagePlugin>) {
688 let chain = self
689 .chains
690 .entry(PluginPositionKey(position))
691 .or_insert_with(|| PluginChain::new(position));
692 chain.push(plugin);
693 }
694
695 /// Register a pre-inference hook by name.
696 ///
697 /// Pre-hooks are called in registration order immediately before inference.
698 /// If any hook returns `Err`, the pipeline aborts with that error message.
699 ///
700 /// # Panics
701 ///
702 /// This function does not panic.
703 pub fn register_pre_hook(&mut self, hook: Arc<dyn InferenceHook>) {
704 let name = hook.name().to_string();
705 debug!(hook = %name, "registered pre-inference hook");
706 self.pre_hooks.push((name, hook));
707 }
708
709 /// Register a post-inference hook by name.
710 ///
711 /// Post-hooks are called in registration order immediately after a
712 /// successful inference. Errors are logged but do not abort the response.
713 ///
714 /// # Panics
715 ///
716 /// This function does not panic.
717 pub fn register_post_hook(&mut self, hook: Arc<dyn InferenceHook>) {
718 let name = hook.name().to_string();
719 debug!(hook = %name, "registered post-inference hook");
720 self.post_hooks.push((name, hook));
721 }
722
723 /// Unregister a plugin by name from the pre-hook list.
724 ///
725 /// Removes the first matching hook. No-op if the name is not found.
726 ///
727 /// # Panics
728 ///
729 /// This function does not panic.
730 pub fn unregister_pre_hook(&mut self, name: &str) {
731 self.pre_hooks.retain(|(n, _)| n != name);
732 }
733
734 /// Unregister a plugin by name from the post-hook list.
735 ///
736 /// Removes the first matching hook. No-op if the name is not found.
737 ///
738 /// # Panics
739 ///
740 /// This function does not panic.
741 pub fn unregister_post_hook(&mut self, name: &str) {
742 self.post_hooks.retain(|(n, _)| n != name);
743 }
744
745 /// Run all registered pre-inference hooks in order.
746 ///
747 /// Returns `Ok(())` if all hooks pass. Returns `Err(message)` at the
748 /// first hook that signals a failure; remaining hooks are not called.
749 ///
750 /// Each hook is run with a per-hook timeout of `hook_timeout` to prevent
751 /// a slow hook from blocking the pipeline indefinitely.
752 ///
753 /// # Errors
754 ///
755 /// Returns the error message from the first failing hook.
756 ///
757 /// # Panics
758 ///
759 /// This function does not panic.
760 pub async fn run_pre_hooks(
761 &self,
762 ctx: &PluginContext,
763 hook_timeout: Duration,
764 ) -> Result<(), String> {
765 for (name, hook) in &self.pre_hooks {
766 let result = tokio::time::timeout(hook_timeout, hook.call(ctx)).await;
767 match result {
768 Ok(Ok(())) => {
769 debug!(hook = %name, "pre-inference hook passed");
770 }
771 Ok(Err(msg)) => {
772 warn!(hook = %name, error = %msg, "pre-inference hook rejected request");
773 return Err(msg);
774 }
775 Err(_elapsed) => {
776 let msg = format!("pre-inference hook '{name}' timed out");
777 warn!(hook = %name, "pre-inference hook timed out");
778 return Err(msg);
779 }
780 }
781 }
782 Ok(())
783 }
784
785 /// Run all registered post-inference hooks in order.
786 ///
787 /// Unlike pre-hooks, post-hook errors are logged but do not abort the
788 /// caller — the response has already been produced. Each hook still
789 /// receives the full [`PluginContext`] (including `response_payload`).
790 ///
791 /// Each hook is run with a per-hook timeout of `hook_timeout`.
792 ///
793 /// # Panics
794 ///
795 /// This function does not panic.
796 pub async fn run_post_hooks(&self, ctx: &PluginContext, hook_timeout: Duration) {
797 for (name, hook) in &self.post_hooks {
798 let result = tokio::time::timeout(hook_timeout, hook.call(ctx)).await;
799 match result {
800 Ok(Ok(())) => {
801 debug!(hook = %name, "post-inference hook completed");
802 }
803 Ok(Err(msg)) => {
804 warn!(hook = %name, error = %msg, "post-inference hook reported error (non-fatal)");
805 }
806 Err(_elapsed) => {
807 warn!(hook = %name, "post-inference hook timed out (non-fatal)");
808 }
809 }
810 }
811 }
812
813 /// Remove all plugins registered at `position`.
814 ///
815 /// After this call, [`chain_for`](Self::chain_for) will return an empty
816 /// chain for that position.
817 ///
818 /// # Panics
819 ///
820 /// This function does not panic.
821 pub fn remove_all(&mut self, position: PluginPosition) {
822 self.chains.remove(&PluginPositionKey(position));
823 }
824
825 /// Return the number of plugins registered at `position`.
826 ///
827 /// # Panics
828 ///
829 /// This function does not panic.
830 pub fn plugin_count(&self, position: PluginPosition) -> usize {
831 self.chains
832 .get(&PluginPositionKey(position))
833 .map_or(0, |c| c.len())
834 }
835
836 /// Return a [`PluginChain`] for the given position.
837 ///
838 /// If no plugins are registered at `position`, the returned chain is
839 /// empty and [`PluginChain::run`] will be a no-op passthrough.
840 ///
841 /// The returned chain is a **clone** — callers may cache it and call
842 /// `run` concurrently without synchronisation.
843 ///
844 /// # Panics
845 ///
846 /// This function does not panic.
847 pub fn chain_for(&self, position: PluginPosition) -> PluginChain {
848 self.chains
849 .get(&PluginPositionKey(position))
850 .cloned()
851 .unwrap_or_else(|| PluginChain::new(position))
852 }
853
854 /// Return `true` if any plugins are registered at `position`.
855 ///
856 /// # Panics
857 ///
858 /// This function does not panic.
859 pub fn has_plugins(&self, position: PluginPosition) -> bool {
860 self.plugin_count(position) > 0
861 }
862
863 /// Return the total number of plugins across all positions.
864 ///
865 /// # Panics
866 ///
867 /// This function does not panic.
868 pub fn total_plugin_count(&self) -> usize {
869 self.chains.values().map(|c| c.len()).sum()
870 }
871
872 /// Return a `Vec` of `(position_label, plugin_count)` for every registered
873 /// position, sorted by label. Useful for diagnostics and logging.
874 ///
875 /// # Panics
876 ///
877 /// This function does not panic.
878 pub fn summary(&self) -> Vec<(String, usize)> {
879 let mut pairs: Vec<(String, usize)> = self
880 .chains
881 .iter()
882 .map(|(k, c)| (k.0.label(), c.len()))
883 .collect();
884 pairs.sort_by(|a, b| a.0.cmp(&b.0));
885 pairs
886 }
887
888 /// Return the number of registered pre-inference hooks.
889 ///
890 /// # Panics
891 ///
892 /// This function does not panic.
893 pub fn pre_hook_count(&self) -> usize {
894 self.pre_hooks.len()
895 }
896
897 /// Return the number of registered post-inference hooks.
898 ///
899 /// # Panics
900 ///
901 /// This function does not panic.
902 pub fn post_hook_count(&self) -> usize {
903 self.post_hooks.len()
904 }
905}
906
907// ============================================================================
908// Round-7 Plugin system: Plugin trait, PluginV2Chain, PluginV2Registry,
909// PluginError, PluginInfo, and built-in plugins.
910// ============================================================================
911
912use crate::PromptRequest;
913
914/// Error returned by [`Plugin`] hooks.
915#[derive(Debug, thiserror::Error)]
916pub enum PluginError {
917 /// The plugin determined that the request should be rejected outright.
918 #[error("request rejected by plugin: {reason}")]
919 RequestRejected {
920 /// Human-readable explanation of why the request was rejected.
921 reason: String,
922 },
923 /// The plugin modified the response (informational, not fatal).
924 #[error("response modified by plugin")]
925 ResponseModified,
926 /// An unrecoverable error occurred inside the plugin.
927 #[error("plugin fatal error: {0}")]
928 Fatal(String),
929}
930
931/// Snapshot of per-plugin call statistics.
932#[derive(Debug, Clone)]
933pub struct PluginInfo {
934 /// Plugin name as returned by [`Plugin::name`].
935 pub name: String,
936 /// Plugin version as returned by [`Plugin::version`].
937 pub version: String,
938 /// Whether the plugin is currently enabled.
939 pub enabled: bool,
940 /// Number of times `on_request` was invoked.
941 pub request_calls: u64,
942 /// Number of times `on_response` was invoked.
943 pub response_calls: u64,
944 /// Total number of errors returned from either hook.
945 pub errors: u64,
946}
947
948/// Core trait for request/response plugins.
949///
950/// Implement this trait and register it with a [`PluginV2Registry`] to
951/// intercept requests before inference and responses after inference.
952pub trait Plugin: Send + Sync {
953 /// Unique human-readable name, e.g. `"profanity-filter"`.
954 fn name(&self) -> &str;
955 /// SemVer string, e.g. `"1.0.0"`.
956 fn version(&self) -> &str;
957 /// Inspect and/or mutate the request before it reaches inference.
958 ///
959 /// Return `Err(PluginError::RequestRejected { .. })` to abort the request.
960 fn on_request(&self, req: &mut PromptRequest) -> Result<(), PluginError>;
961 /// Inspect and/or mutate the response tokens after inference.
962 ///
963 /// Return `Err(PluginError::ResponseModified)` or `Err(PluginError::Fatal)`
964 /// to signal post-processing issues (pipeline may log and continue).
965 fn on_response(&self, resp: &mut Vec<String>) -> Result<(), PluginError>;
966}
967
968/// Internal entry in [`PluginV2Registry`].
969struct PluginEntry {
970 plugin: Box<dyn Plugin>,
971 enabled: bool,
972 request_calls: u64,
973 response_calls: u64,
974 errors: u64,
975}
976
977/// Ordered list of [`Plugin`]s executed sequentially on each request/response.
978///
979/// Plugins run in insertion order. If any `on_request` hook returns
980/// `Err(PluginError::RequestRejected)`, the chain stops and the error is
981/// propagated. `on_response` errors are collected but do not short-circuit
982/// the remaining response plugins.
983#[derive(Default)]
984pub struct PluginV2Chain {
985 plugins: Vec<Box<dyn Plugin>>,
986}
987
988impl PluginV2Chain {
989 /// Create an empty chain.
990 #[must_use]
991 pub fn new() -> Self {
992 Self::default()
993 }
994
995 /// Append a plugin to the end of the chain.
996 pub fn push(&mut self, plugin: Box<dyn Plugin>) {
997 self.plugins.push(plugin);
998 }
999
1000 /// Run `on_request` for all plugins in order.
1001 ///
1002 /// Stops at the first `RequestRejected` error.
1003 pub fn run_request(&self, req: &mut PromptRequest) -> Result<(), PluginError> {
1004 for plugin in &self.plugins {
1005 plugin.on_request(req)?;
1006 }
1007 Ok(())
1008 }
1009
1010 /// Run `on_response` for all plugins in order.
1011 ///
1012 /// Continues even if individual plugins return an error (errors are
1013 /// returned at the end as the last encountered error).
1014 pub fn run_response(&self, resp: &mut Vec<String>) -> Result<(), PluginError> {
1015 let mut last_err: Option<PluginError> = None;
1016 for plugin in &self.plugins {
1017 if let Err(e) = plugin.on_response(resp) {
1018 last_err = Some(e);
1019 }
1020 }
1021 match last_err {
1022 None => Ok(()),
1023 Some(e) => Err(e),
1024 }
1025 }
1026
1027 /// Return the number of plugins in this chain.
1028 pub fn len(&self) -> usize {
1029 self.plugins.len()
1030 }
1031
1032 /// Return `true` if no plugins are registered.
1033 pub fn is_empty(&self) -> bool {
1034 self.plugins.is_empty()
1035 }
1036}
1037
1038/// Runtime registry of [`Plugin`] instances with enable/disable support and
1039/// per-plugin call statistics.
1040#[derive(Default)]
1041pub struct PluginV2Registry {
1042 entries: Vec<PluginEntry>,
1043}
1044
1045impl PluginV2Registry {
1046 /// Create a new, empty registry.
1047 #[must_use]
1048 pub fn new() -> Self {
1049 Self::default()
1050 }
1051
1052 /// Register a plugin. Plugins execute in registration order.
1053 pub fn register(&mut self, plugin: Box<dyn Plugin>) {
1054 self.entries.push(PluginEntry {
1055 plugin,
1056 enabled: true,
1057 request_calls: 0,
1058 response_calls: 0,
1059 errors: 0,
1060 });
1061 }
1062
1063 /// Disable a plugin by name. Disabled plugins are skipped during
1064 /// `run_request_hooks` and `run_response_hooks` but remain registered.
1065 ///
1066 /// No-op if the name is not found.
1067 pub fn disable(&mut self, name: &str) {
1068 for entry in &mut self.entries {
1069 if entry.plugin.name() == name {
1070 entry.enabled = false;
1071 }
1072 }
1073 }
1074
1075 /// Re-enable a plugin that was previously disabled.
1076 ///
1077 /// No-op if the name is not found.
1078 pub fn enable(&mut self, name: &str) {
1079 for entry in &mut self.entries {
1080 if entry.plugin.name() == name {
1081 entry.enabled = true;
1082 }
1083 }
1084 }
1085
1086 /// List all registered plugins with their current statistics.
1087 pub fn list(&self) -> Vec<PluginInfo> {
1088 self.entries
1089 .iter()
1090 .map(|e| PluginInfo {
1091 name: e.plugin.name().to_string(),
1092 version: e.plugin.version().to_string(),
1093 enabled: e.enabled,
1094 request_calls: e.request_calls,
1095 response_calls: e.response_calls,
1096 errors: e.errors,
1097 })
1098 .collect()
1099 }
1100
1101 /// Run `on_request` for all enabled plugins in registration order.
1102 ///
1103 /// Stops at the first `RequestRejected` error, records the error in stats.
1104 pub fn run_request_hooks(&mut self, req: &mut PromptRequest) -> Result<(), PluginError> {
1105 for entry in &mut self.entries {
1106 if !entry.enabled {
1107 continue;
1108 }
1109 entry.request_calls += 1;
1110 if let Err(e) = entry.plugin.on_request(req) {
1111 entry.errors += 1;
1112 return Err(e);
1113 }
1114 }
1115 Ok(())
1116 }
1117
1118 /// Run `on_response` for all enabled plugins in registration order.
1119 ///
1120 /// Does not short-circuit on error — all enabled plugins run. The last
1121 /// error encountered is returned, if any.
1122 pub fn run_response_hooks(&mut self, resp: &mut Vec<String>) -> Result<(), PluginError> {
1123 let mut last_err: Option<PluginError> = None;
1124 for entry in &mut self.entries {
1125 if !entry.enabled {
1126 continue;
1127 }
1128 entry.response_calls += 1;
1129 if let Err(e) = entry.plugin.on_response(resp) {
1130 entry.errors += 1;
1131 last_err = Some(e);
1132 }
1133 }
1134 match last_err {
1135 None => Ok(()),
1136 Some(e) => Err(e),
1137 }
1138 }
1139}
1140
1141// ============================================================================
1142// Built-in plugins
1143// ============================================================================
1144
1145/// A plugin that blocks requests containing any word from a configurable list.
1146///
1147/// The check is case-insensitive.
1148pub struct ProfanityFilterPlugin {
1149 /// Words that, if present in the prompt, cause the request to be rejected.
1150 pub word_list: Vec<String>,
1151}
1152
1153impl ProfanityFilterPlugin {
1154 /// Create a new filter with the given list of blocked words.
1155 pub fn new(words: impl IntoIterator<Item = impl Into<String>>) -> Self {
1156 Self {
1157 word_list: words.into_iter().map(Into::into).collect(),
1158 }
1159 }
1160}
1161
1162impl Plugin for ProfanityFilterPlugin {
1163 fn name(&self) -> &str {
1164 "profanity-filter"
1165 }
1166 fn version(&self) -> &str {
1167 "1.0.0"
1168 }
1169 fn on_request(&self, req: &mut PromptRequest) -> Result<(), PluginError> {
1170 let lower = req.input.to_lowercase();
1171 for word in &self.word_list {
1172 if lower.contains(word.to_lowercase().as_str()) {
1173 return Err(PluginError::RequestRejected {
1174 reason: format!("prompt contains blocked word: {word}"),
1175 });
1176 }
1177 }
1178 Ok(())
1179 }
1180 fn on_response(&self, _resp: &mut Vec<String>) -> Result<(), PluginError> {
1181 Ok(())
1182 }
1183}
1184
1185/// A plugin that truncates response token lists to at most `max_tokens` items.
1186pub struct ResponseLengthCapPlugin {
1187 /// Maximum number of token strings to keep in the response.
1188 pub max_tokens: usize,
1189}
1190
1191impl ResponseLengthCapPlugin {
1192 /// Create a new cap plugin.
1193 pub fn new(max_tokens: usize) -> Self {
1194 Self { max_tokens }
1195 }
1196}
1197
1198impl Plugin for ResponseLengthCapPlugin {
1199 fn name(&self) -> &str {
1200 "response-length-cap"
1201 }
1202 fn version(&self) -> &str {
1203 "1.0.0"
1204 }
1205 fn on_request(&self, _req: &mut PromptRequest) -> Result<(), PluginError> {
1206 Ok(())
1207 }
1208 fn on_response(&self, resp: &mut Vec<String>) -> Result<(), PluginError> {
1209 if resp.len() > self.max_tokens {
1210 resp.truncate(self.max_tokens);
1211 return Err(PluginError::ResponseModified);
1212 }
1213 Ok(())
1214 }
1215}
1216
1217/// A plugin that records per-request start/end timestamps (milliseconds since
1218/// Unix epoch) into a shared latency log.
1219///
1220/// Call `LatencyLoggerPlugin::record_start` before inference and
1221/// `LatencyLoggerPlugin::record_end` after inference to append the elapsed
1222/// duration to the log.
1223pub struct LatencyLoggerPlugin {
1224 /// Accumulated latency samples in milliseconds.
1225 pub log: std::sync::Mutex<Vec<u64>>,
1226}
1227
1228impl LatencyLoggerPlugin {
1229 /// Create a new logger with an empty latency log.
1230 pub fn new() -> Self {
1231 Self {
1232 log: std::sync::Mutex::new(Vec::new()),
1233 }
1234 }
1235
1236 /// Record a completed request's latency in milliseconds.
1237 pub fn record_latency(&self, latency_ms: u64) {
1238 if let Ok(mut guard) = self.log.lock() {
1239 guard.push(latency_ms);
1240 }
1241 }
1242
1243 /// Return a copy of all recorded latencies.
1244 pub fn samples(&self) -> Vec<u64> {
1245 self.log.lock().map(|g| g.clone()).unwrap_or_default()
1246 }
1247}
1248
1249impl Default for LatencyLoggerPlugin {
1250 fn default() -> Self {
1251 Self::new()
1252 }
1253}
1254
1255impl Plugin for LatencyLoggerPlugin {
1256 fn name(&self) -> &str {
1257 "latency-logger"
1258 }
1259 fn version(&self) -> &str {
1260 "1.0.0"
1261 }
1262 fn on_request(&self, _req: &mut PromptRequest) -> Result<(), PluginError> {
1263 // Latency is recorded externally via record_latency().
1264 Ok(())
1265 }
1266 fn on_response(&self, _resp: &mut Vec<String>) -> Result<(), PluginError> {
1267 Ok(())
1268 }
1269}
1270
1271// ============================================================================
1272// Tests — Round 7 Plugin system (20+ unit tests)
1273// ============================================================================
1274
1275#[cfg(test)]
1276mod plugin_v2_tests {
1277 use super::*;
1278 use std::collections::HashMap;
1279
1280 fn make_req(input: &str) -> PromptRequest {
1281 PromptRequest {
1282 session: crate::SessionId::new("s1"),
1283 request_id: "r1".to_string(),
1284 input: input.to_string(),
1285 meta: HashMap::new(),
1286 deadline: None,
1287 }
1288 }
1289
1290 // ---- ProfanityFilterPlugin ----
1291
1292 #[test]
1293 fn profanity_filter_rejects_blocked_word() {
1294 let p = ProfanityFilterPlugin::new(["badword"]);
1295 let mut req = make_req("this has badword in it");
1296 assert!(matches!(
1297 p.on_request(&mut req),
1298 Err(PluginError::RequestRejected { .. })
1299 ));
1300 }
1301
1302 #[test]
1303 fn profanity_filter_case_insensitive() {
1304 let p = ProfanityFilterPlugin::new(["BadWord"]);
1305 let mut req = make_req("BADWORD is uppercase");
1306 assert!(matches!(
1307 p.on_request(&mut req),
1308 Err(PluginError::RequestRejected { .. })
1309 ));
1310 }
1311
1312 #[test]
1313 fn profanity_filter_passes_clean_request() {
1314 let p = ProfanityFilterPlugin::new(["blocked"]);
1315 let mut req = make_req("this is totally fine");
1316 assert!(p.on_request(&mut req).is_ok());
1317 }
1318
1319 #[test]
1320 fn profanity_filter_empty_word_list_passes() {
1321 let p = ProfanityFilterPlugin::new([] as [&str; 0]);
1322 let mut req = make_req("anything goes");
1323 assert!(p.on_request(&mut req).is_ok());
1324 }
1325
1326 #[test]
1327 fn profanity_filter_on_response_is_noop() {
1328 let p = ProfanityFilterPlugin::new(["x"]);
1329 let mut resp = vec!["a".to_string()];
1330 assert!(p.on_response(&mut resp).is_ok());
1331 }
1332
1333 // ---- ResponseLengthCapPlugin ----
1334
1335 #[test]
1336 fn length_cap_truncates_over_limit() {
1337 let p = ResponseLengthCapPlugin::new(3);
1338 let mut resp: Vec<String> = (0..5).map(|i| i.to_string()).collect();
1339 let result = p.on_response(&mut resp);
1340 assert!(matches!(result, Err(PluginError::ResponseModified)));
1341 assert_eq!(resp.len(), 3);
1342 }
1343
1344 #[test]
1345 fn length_cap_passes_exact_limit() {
1346 let p = ResponseLengthCapPlugin::new(3);
1347 let mut resp: Vec<String> = (0..3).map(|i| i.to_string()).collect();
1348 assert!(p.on_response(&mut resp).is_ok());
1349 assert_eq!(resp.len(), 3);
1350 }
1351
1352 #[test]
1353 fn length_cap_passes_under_limit() {
1354 let p = ResponseLengthCapPlugin::new(10);
1355 let mut resp = vec!["a".to_string(), "b".to_string()];
1356 assert!(p.on_response(&mut resp).is_ok());
1357 }
1358
1359 #[test]
1360 fn length_cap_on_request_is_noop() {
1361 let p = ResponseLengthCapPlugin::new(1);
1362 let mut req = make_req("hello");
1363 assert!(p.on_request(&mut req).is_ok());
1364 }
1365
1366 // ---- LatencyLoggerPlugin ----
1367
1368 #[test]
1369 fn latency_logger_records_samples() {
1370 let p = LatencyLoggerPlugin::new();
1371 p.record_latency(42);
1372 p.record_latency(99);
1373 let samples = p.samples();
1374 assert_eq!(samples, vec![42, 99]);
1375 }
1376
1377 #[test]
1378 fn latency_logger_on_request_ok() {
1379 let p = LatencyLoggerPlugin::new();
1380 let mut req = make_req("hello");
1381 assert!(p.on_request(&mut req).is_ok());
1382 }
1383
1384 #[test]
1385 fn latency_logger_on_response_ok() {
1386 let p = LatencyLoggerPlugin::new();
1387 let mut resp = vec!["tok".to_string()];
1388 assert!(p.on_response(&mut resp).is_ok());
1389 }
1390
1391 // ---- PluginV2Registry ----
1392
1393 #[test]
1394 fn registry_list_returns_all() {
1395 let mut reg = PluginV2Registry::new();
1396 reg.register(Box::new(ProfanityFilterPlugin::new(["x"])));
1397 reg.register(Box::new(ResponseLengthCapPlugin::new(10)));
1398 let info = reg.list();
1399 assert_eq!(info.len(), 2);
1400 assert_eq!(info[0].name, "profanity-filter");
1401 assert_eq!(info[1].name, "response-length-cap");
1402 }
1403
1404 #[test]
1405 fn registry_disable_skips_plugin() {
1406 let mut reg = PluginV2Registry::new();
1407 reg.register(Box::new(ProfanityFilterPlugin::new(["bad"])));
1408 reg.disable("profanity-filter");
1409 let mut req = make_req("bad content here");
1410 // Disabled plugin should not reject.
1411 assert!(reg.run_request_hooks(&mut req).is_ok());
1412 }
1413
1414 #[test]
1415 fn registry_enable_reenables_plugin() {
1416 let mut reg = PluginV2Registry::new();
1417 reg.register(Box::new(ProfanityFilterPlugin::new(["bad"])));
1418 reg.disable("profanity-filter");
1419 reg.enable("profanity-filter");
1420 let mut req = make_req("bad content");
1421 assert!(matches!(
1422 reg.run_request_hooks(&mut req),
1423 Err(PluginError::RequestRejected { .. })
1424 ));
1425 }
1426
1427 #[test]
1428 fn registry_tracks_request_calls() {
1429 let mut reg = PluginV2Registry::new();
1430 reg.register(Box::new(ProfanityFilterPlugin::new([] as [&str; 0])));
1431 let mut req = make_req("hello");
1432 reg.run_request_hooks(&mut req).ok();
1433 reg.run_request_hooks(&mut req).ok();
1434 let info = reg.list();
1435 assert_eq!(info[0].request_calls, 2);
1436 }
1437
1438 #[test]
1439 fn registry_tracks_response_calls() {
1440 let mut reg = PluginV2Registry::new();
1441 reg.register(Box::new(ResponseLengthCapPlugin::new(100)));
1442 let mut resp = vec!["tok".to_string()];
1443 reg.run_response_hooks(&mut resp).ok();
1444 let info = reg.list();
1445 assert_eq!(info[0].response_calls, 1);
1446 }
1447
1448 #[test]
1449 fn registry_tracks_errors() {
1450 let mut reg = PluginV2Registry::new();
1451 reg.register(Box::new(ProfanityFilterPlugin::new(["bad"])));
1452 let mut req = make_req("bad");
1453 reg.run_request_hooks(&mut req).ok();
1454 let info = reg.list();
1455 assert_eq!(info[0].errors, 1);
1456 }
1457
1458 #[test]
1459 fn registry_run_response_hooks_all_plugins_run_despite_error() {
1460 let mut reg = PluginV2Registry::new();
1461 reg.register(Box::new(ResponseLengthCapPlugin::new(0)));
1462 reg.register(Box::new(ResponseLengthCapPlugin::new(100)));
1463 let mut resp = vec!["tok1".to_string(), "tok2".to_string()];
1464 // First plugin truncates and returns Err; second should still run.
1465 let _ = reg.run_response_hooks(&mut resp);
1466 let info = reg.list();
1467 assert_eq!(info[0].response_calls, 1);
1468 assert_eq!(info[1].response_calls, 1);
1469 }
1470
1471 #[test]
1472 fn registry_empty_runs_ok() {
1473 let mut reg = PluginV2Registry::new();
1474 let mut req = make_req("hello");
1475 assert!(reg.run_request_hooks(&mut req).is_ok());
1476 let mut resp = vec![];
1477 assert!(reg.run_response_hooks(&mut resp).is_ok());
1478 }
1479
1480 #[test]
1481 fn plugin_info_enabled_flag() {
1482 let mut reg = PluginV2Registry::new();
1483 reg.register(Box::new(LatencyLoggerPlugin::new()));
1484 reg.disable("latency-logger");
1485 let info = reg.list();
1486 assert!(!info[0].enabled);
1487 reg.enable("latency-logger");
1488 let info = reg.list();
1489 assert!(info[0].enabled);
1490 }
1491
1492 #[test]
1493 fn plugin_chain_v2_run_request_all_pass() {
1494 let mut chain = PluginV2Chain::new();
1495 chain.push(Box::new(ProfanityFilterPlugin::new([] as [&str; 0])));
1496 chain.push(Box::new(ResponseLengthCapPlugin::new(100)));
1497 let mut req = make_req("clean request");
1498 assert!(chain.run_request(&mut req).is_ok());
1499 }
1500
1501 #[test]
1502 fn plugin_chain_v2_run_request_stops_on_rejection() {
1503 let mut chain = PluginV2Chain::new();
1504 chain.push(Box::new(ProfanityFilterPlugin::new(["bad"])));
1505 chain.push(Box::new(ProfanityFilterPlugin::new(["other"])));
1506 let mut req = make_req("bad content");
1507 assert!(matches!(
1508 chain.run_request(&mut req),
1509 Err(PluginError::RequestRejected { .. })
1510 ));
1511 }
1512
1513 #[test]
1514 fn plugin_chain_v2_len_is_empty() {
1515 let chain = PluginV2Chain::new();
1516 assert!(chain.is_empty());
1517 assert_eq!(chain.len(), 0);
1518 }
1519
1520 #[test]
1521 fn plugin_error_display() {
1522 let e = PluginError::RequestRejected {
1523 reason: "bad".to_string(),
1524 };
1525 assert!(e.to_string().contains("bad"));
1526 let e2 = PluginError::ResponseModified;
1527 assert!(e2.to_string().contains("modified"));
1528 let e3 = PluginError::Fatal("crash".to_string());
1529 assert!(e3.to_string().contains("crash"));
1530 }
1531}
1532
1533// ============================================================================
1534// Tests
1535// ============================================================================
1536
1537#[cfg(test)]
1538mod tests {
1539 use super::*;
1540
1541 struct PassthroughPlugin;
1542
1543 #[async_trait]
1544 impl StagePlugin for PassthroughPlugin {
1545 fn name(&self) -> &'static str {
1546 "passthrough"
1547 }
1548 async fn process(&self, input: PluginInput) -> PluginOutput {
1549 PluginOutput::passthrough(input)
1550 }
1551 }
1552
1553 struct AbortPlugin;
1554
1555 #[async_trait]
1556 impl StagePlugin for AbortPlugin {
1557 fn name(&self) -> &'static str {
1558 "abort"
1559 }
1560 async fn process(&self, input: PluginInput) -> PluginOutput {
1561 PluginOutput::abort(input)
1562 }
1563 }
1564
1565 struct MetaMutatorPlugin;
1566
1567 #[async_trait]
1568 impl StagePlugin for MetaMutatorPlugin {
1569 fn name(&self) -> &'static str {
1570 "meta-mutator"
1571 }
1572 async fn process(&self, mut input: PluginInput) -> PluginOutput {
1573 input.metadata.insert("mutated".to_string(), "yes".to_string());
1574 PluginOutput::passthrough(input)
1575 }
1576 }
1577
1578 fn make_input() -> PluginInput {
1579 PluginInput {
1580 request_id: "req-test".to_string(),
1581 session_id: "session-test".to_string(),
1582 payload: serde_json::json!({"prompt": "hello"}),
1583 metadata: HashMap::new(),
1584 }
1585 }
1586
1587 #[tokio::test]
1588 async fn empty_chain_is_passthrough() {
1589 let chain = PluginChain::new(PluginPosition::Before(PipelineStage::Inference));
1590 let input = make_input();
1591 let output = chain.run(input).await;
1592 assert_eq!(output.status, PluginStatus::Continue);
1593 }
1594
1595 #[tokio::test]
1596 async fn single_passthrough_plugin() {
1597 let mut chain = PluginChain::new(PluginPosition::Before(PipelineStage::Rag));
1598 chain.push(Arc::new(PassthroughPlugin));
1599 let output = chain.run(make_input()).await;
1600 assert_eq!(output.status, PluginStatus::Continue);
1601 }
1602
1603 #[tokio::test]
1604 async fn abort_plugin_stops_chain() {
1605 let mut chain = PluginChain::new(PluginPosition::After(PipelineStage::Inference));
1606 chain.push(Arc::new(AbortPlugin));
1607 chain.push(Arc::new(PassthroughPlugin)); // should NOT run
1608 let output = chain.run(make_input()).await;
1609 assert_eq!(output.status, PluginStatus::Abort);
1610 }
1611
1612 #[tokio::test]
1613 async fn meta_mutator_modifies_metadata() {
1614 let mut chain = PluginChain::new(PluginPosition::Before(PipelineStage::Post));
1615 chain.push(Arc::new(MetaMutatorPlugin));
1616 let output = chain.run(make_input()).await;
1617 assert_eq!(
1618 output.input.metadata.get("mutated").map(String::as_str),
1619 Some("yes")
1620 );
1621 }
1622
1623 #[tokio::test]
1624 async fn registry_counts_correctly() {
1625 let mut registry = PluginRegistry::new();
1626 let pos = PluginPosition::Before(PipelineStage::Inference);
1627 registry.register(pos, Arc::new(PassthroughPlugin));
1628 registry.register(pos, Arc::new(PassthroughPlugin));
1629 assert_eq!(registry.plugin_count(pos), 2);
1630 assert_eq!(registry.total_plugin_count(), 2);
1631 }
1632
1633 #[tokio::test]
1634 async fn registry_remove_all_clears_chain() {
1635 let mut registry = PluginRegistry::new();
1636 let pos = PluginPosition::After(PipelineStage::Rag);
1637 registry.register(pos, Arc::new(PassthroughPlugin));
1638 registry.remove_all(pos);
1639 assert_eq!(registry.plugin_count(pos), 0);
1640 // chain_for on an empty position is a no-op passthrough
1641 let chain = registry.chain_for(pos);
1642 let output = chain.run(make_input()).await;
1643 assert_eq!(output.status, PluginStatus::Continue);
1644 }
1645
1646 #[tokio::test]
1647 async fn registry_summary_sorted() {
1648 let mut registry = PluginRegistry::new();
1649 registry.register(
1650 PluginPosition::Before(PipelineStage::Stream),
1651 Arc::new(PassthroughPlugin),
1652 );
1653 registry.register(
1654 PluginPosition::After(PipelineStage::Rag),
1655 Arc::new(PassthroughPlugin),
1656 );
1657 let summary = registry.summary();
1658 // Should be sorted alphabetically: "after:rag" < "before:stream"
1659 assert_eq!(summary[0].0, "after:rag");
1660 assert_eq!(summary[1].0, "before:stream");
1661 }
1662
1663 #[test]
1664 fn plugin_position_label() {
1665 assert_eq!(
1666 PluginPosition::Before(PipelineStage::Inference).label(),
1667 "before:inference"
1668 );
1669 assert_eq!(
1670 PluginPosition::After(PipelineStage::Post).label(),
1671 "after:post"
1672 );
1673 }
1674
1675 #[test]
1676 fn plugin_output_constructors() {
1677 let input = make_input();
1678 let out = PluginOutput::passthrough(input.clone());
1679 assert_eq!(out.status, PluginStatus::Continue);
1680
1681 let out = PluginOutput::abort(input.clone());
1682 assert_eq!(out.status, PluginStatus::Abort);
1683
1684 let out = PluginOutput::error(input.clone(), "oops");
1685 assert!(matches!(out.status, PluginStatus::Error(ref s) if s == "oops"));
1686
1687 let modified = PluginOutput::modified(input, serde_json::json!({"x": 1}));
1688 assert_eq!(modified.input.payload, serde_json::json!({"x": 1}));
1689 assert_eq!(modified.status, PluginStatus::Continue);
1690 }
1691}