Skip to main content

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}