[model-gateway] Tighten visibility across data_connector and grpc module (#16516)
This commit is contained in:
@@ -2,28 +2,28 @@ use std::collections::HashMap;
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
pub fn parse_tool_calls(raw: Option<String>) -> Result<Vec<Value>, String> {
|
||||
pub(super) fn parse_tool_calls(raw: Option<String>) -> Result<Vec<Value>, String> {
|
||||
match raw {
|
||||
Some(s) if !s.is_empty() => serde_json::from_str(&s).map_err(|e| e.to_string()),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn parse_metadata(raw: Option<String>) -> Result<HashMap<String, Value>, String> {
|
||||
pub(super) fn parse_metadata(raw: Option<String>) -> Result<HashMap<String, Value>, String> {
|
||||
match raw {
|
||||
Some(s) if !s.is_empty() => serde_json::from_str(&s).map_err(|e| e.to_string()),
|
||||
_ => Ok(HashMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn parse_raw_response(raw: Option<String>) -> Result<Value, String> {
|
||||
pub(super) fn parse_raw_response(raw: Option<String>) -> Result<Value, String> {
|
||||
match raw {
|
||||
Some(s) if !s.is_empty() => serde_json::from_str(&s).map_err(|e| e.to_string()),
|
||||
_ => Ok(Value::Null),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn parse_json_value(raw: Option<String>) -> Result<Value, String> {
|
||||
pub(super) fn parse_json_value(raw: Option<String>) -> Result<Value, String> {
|
||||
match raw {
|
||||
Some(s) if !s.is_empty() => serde_json::from_str(&s).map_err(|e| e.to_string()),
|
||||
_ => Ok(Value::Array(vec![])),
|
||||
|
||||
@@ -288,7 +288,8 @@ impl MemoryResponseStorage {
|
||||
}
|
||||
|
||||
/// Get statistics about the store
|
||||
pub fn stats(&self) -> MemoryStoreStats {
|
||||
#[allow(dead_code)]
|
||||
pub(super) fn stats(&self) -> MemoryStoreStats {
|
||||
let store = self.store.read();
|
||||
MemoryStoreStats {
|
||||
response_count: store.responses.len(),
|
||||
@@ -459,7 +460,8 @@ impl ResponseStorage for MemoryResponseStorage {
|
||||
|
||||
/// Statistics for the memory store
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MemoryStoreStats {
|
||||
#[allow(dead_code)]
|
||||
pub(super) struct MemoryStoreStats {
|
||||
pub response_count: usize,
|
||||
pub identifier_count: usize,
|
||||
}
|
||||
|
||||
@@ -18,7 +18,7 @@ use super::core::*;
|
||||
|
||||
/// No-op implementation that synthesizes conversation responses without persistence
|
||||
#[derive(Default, Debug, Clone)]
|
||||
pub struct NoOpConversationStorage;
|
||||
pub(super) struct NoOpConversationStorage;
|
||||
|
||||
impl NoOpConversationStorage {
|
||||
pub fn new() -> Self {
|
||||
@@ -61,7 +61,7 @@ impl ConversationStorage for NoOpConversationStorage {
|
||||
|
||||
/// No-op conversation item storage (does nothing)
|
||||
#[derive(Clone, Copy, Default)]
|
||||
pub struct NoOpConversationItemStorage;
|
||||
pub(super) struct NoOpConversationItemStorage;
|
||||
|
||||
impl NoOpConversationItemStorage {
|
||||
pub fn new() -> Self {
|
||||
@@ -136,7 +136,7 @@ impl ConversationItemStorage for NoOpConversationItemStorage {
|
||||
// ============================================================================
|
||||
|
||||
/// No-op implementation of response storage (does nothing)
|
||||
pub struct NoOpResponseStorage;
|
||||
pub(super) struct NoOpResponseStorage;
|
||||
|
||||
impl NoOpResponseStorage {
|
||||
pub fn new() -> Self {
|
||||
|
||||
@@ -232,7 +232,7 @@ impl Manager for OracleConnectionManager {
|
||||
// ============================================================================
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct OracleConversationStorage {
|
||||
pub(super) struct OracleConversationStorage {
|
||||
store: OracleStore,
|
||||
}
|
||||
|
||||
@@ -420,7 +420,7 @@ impl ConversationStorage for OracleConversationStorage {
|
||||
// ============================================================================
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct OracleConversationItemStorage {
|
||||
pub(super) struct OracleConversationItemStorage {
|
||||
store: OracleStore,
|
||||
}
|
||||
|
||||
@@ -775,7 +775,7 @@ const SELECT_BASE: &str = "SELECT id, previous_response_id, input, instructions,
|
||||
tool_calls, metadata, created_at, safety_identifier, model, conversation_id, raw_response FROM responses";
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct OracleResponseStorage {
|
||||
pub(super) struct OracleResponseStorage {
|
||||
store: OracleStore,
|
||||
}
|
||||
|
||||
|
||||
@@ -58,7 +58,7 @@ impl Clone for PostgresStore {
|
||||
}
|
||||
}
|
||||
|
||||
pub struct PostgresConversationStorage {
|
||||
pub(super) struct PostgresConversationStorage {
|
||||
store: PostgresStore,
|
||||
}
|
||||
|
||||
@@ -198,7 +198,7 @@ impl ConversationStorage for PostgresConversationStorage {
|
||||
}
|
||||
}
|
||||
|
||||
pub struct PostgresConversationItemStorage {
|
||||
pub(super) struct PostgresConversationItemStorage {
|
||||
store: PostgresStore,
|
||||
}
|
||||
|
||||
@@ -477,7 +477,7 @@ impl ConversationItemStorage for PostgresConversationItemStorage {
|
||||
}
|
||||
}
|
||||
|
||||
pub struct PostgresResponseStorage {
|
||||
pub(super) struct PostgresResponseStorage {
|
||||
store: PostgresStore,
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
//! Shared code for both regular and harmony routers
|
||||
|
||||
pub mod response_collection;
|
||||
pub mod response_formatting;
|
||||
pub mod responses;
|
||||
pub mod stages;
|
||||
pub(crate) mod response_collection;
|
||||
pub(crate) mod response_formatting;
|
||||
pub(crate) mod responses;
|
||||
pub(crate) mod stages;
|
||||
|
||||
@@ -21,7 +21,7 @@ use crate::routers::{
|
||||
///
|
||||
/// # Returns
|
||||
/// Vector of GenerateComplete responses, one per index (n parameter)
|
||||
pub async fn collect_responses(
|
||||
pub(crate) async fn collect_responses(
|
||||
execution_result: ExecutionResult,
|
||||
merge_logprobs: bool,
|
||||
) -> Result<Vec<ProtoGenerateComplete>, Response> {
|
||||
|
||||
@@ -16,7 +16,7 @@ use crate::{protocols::common::Usage, routers::grpc::proto_wrapper::ProtoGenerat
|
||||
///
|
||||
/// # Returns
|
||||
/// Usage object with aggregated token counts
|
||||
pub fn build_usage(responses: &[ProtoGenerateComplete]) -> Usage {
|
||||
pub(crate) fn build_usage(responses: &[ProtoGenerateComplete]) -> Usage {
|
||||
let total_prompt_tokens: u32 = responses.iter().map(|r| r.prompt_tokens() as u32).sum();
|
||||
let total_completion_tokens: u32 = responses.iter().map(|r| r.completion_tokens() as u32).sum();
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ use crate::{
|
||||
///
|
||||
/// Retrieves a stored response from the database.
|
||||
/// Used by both regular and harmony implementations.
|
||||
pub async fn get_response_impl(ctx: &ResponsesContext, response_id: &str) -> Response {
|
||||
pub(crate) async fn get_response_impl(ctx: &ResponsesContext, response_id: &str) -> Response {
|
||||
let resp_id = ResponseId::from(response_id);
|
||||
|
||||
// Retrieve response from storage
|
||||
@@ -38,7 +38,7 @@ pub async fn get_response_impl(ctx: &ResponsesContext, response_id: &str) -> Res
|
||||
/// Implementation for POST /v1/responses/{response_id}/cancel
|
||||
///
|
||||
/// Cancels a background response if it's still in progress.
|
||||
pub async fn cancel_response_impl(ctx: &ResponsesContext, response_id: &str) -> Response {
|
||||
pub(crate) async fn cancel_response_impl(ctx: &ResponsesContext, response_id: &str) -> Response {
|
||||
let resp_id = ResponseId::from(response_id);
|
||||
|
||||
// Retrieve response from storage to check if it exists and get current status
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
//! Shared response functionality used by both regular and harmony implementations
|
||||
|
||||
pub mod handlers;
|
||||
pub mod streaming;
|
||||
pub mod utils;
|
||||
pub(crate) mod handlers;
|
||||
pub(crate) mod streaming;
|
||||
pub(crate) mod utils;
|
||||
|
||||
pub use handlers::{cancel_response_impl, get_response_impl};
|
||||
pub use streaming::{build_sse_response, OutputItemType, ResponseStreamEventEmitter};
|
||||
pub use utils::{ensure_mcp_connection, persist_response_if_needed};
|
||||
// Re-export commonly used items
|
||||
pub(crate) use streaming::build_sse_response;
|
||||
pub(crate) use utils::{ensure_mcp_connection, persist_response_if_needed};
|
||||
|
||||
@@ -25,7 +25,7 @@ use crate::{
|
||||
routers::grpc::harmony::responses::ToolResult,
|
||||
};
|
||||
|
||||
pub enum OutputItemType {
|
||||
pub(crate) enum OutputItemType {
|
||||
Message,
|
||||
McpListTools,
|
||||
McpCall,
|
||||
@@ -67,7 +67,7 @@ struct OutputItemState {
|
||||
/// - response.mcp_call_arguments.done
|
||||
/// - response.mcp_call.completed
|
||||
/// - response.mcp_call.failed
|
||||
pub struct ResponseStreamEventEmitter {
|
||||
pub(crate) struct ResponseStreamEventEmitter {
|
||||
sequence_number: u64,
|
||||
pub response_id: String,
|
||||
model: String,
|
||||
@@ -828,7 +828,9 @@ impl ResponseStreamEventEmitter {
|
||||
/// Build a Server-Sent Events (SSE) response
|
||||
///
|
||||
/// Creates a Response with proper SSE headers and streaming body.
|
||||
pub fn build_sse_response(rx: mpsc::UnboundedReceiver<Result<Bytes, std::io::Error>>) -> Response {
|
||||
pub(crate) fn build_sse_response(
|
||||
rx: mpsc::UnboundedReceiver<Result<Bytes, std::io::Error>>,
|
||||
) -> Response {
|
||||
let stream = UnboundedReceiverStream::new(rx);
|
||||
Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
|
||||
@@ -23,7 +23,7 @@ use crate::{
|
||||
///
|
||||
/// Checks if request declares MCP tools, and if so, validates that
|
||||
/// the MCP client can be created and connected.
|
||||
pub async fn ensure_mcp_connection(
|
||||
pub(crate) async fn ensure_mcp_connection(
|
||||
mcp_manager: &Arc<McpManager>,
|
||||
tools: Option<&[ResponseTool]>,
|
||||
) -> Result<bool, Response> {
|
||||
@@ -56,7 +56,7 @@ pub async fn ensure_mcp_connection(
|
||||
}
|
||||
|
||||
/// Validate that workers are available for the requested model
|
||||
pub fn validate_worker_availability(
|
||||
pub(crate) fn validate_worker_availability(
|
||||
worker_registry: &Arc<WorkerRegistry>,
|
||||
model: &str,
|
||||
) -> Option<Response> {
|
||||
@@ -90,7 +90,7 @@ pub fn validate_worker_availability(
|
||||
/// the initial conversion from ResponsesRequest to ChatCompletionRequest. MCP tools
|
||||
/// are merged later by the tool loop before being sent to the chat pipeline, where
|
||||
/// tool_choice constraints are generated for ALL tools (function + MCP combined).
|
||||
pub fn extract_tools_from_response_tools(
|
||||
pub(crate) fn extract_tools_from_response_tools(
|
||||
response_tools: Option<&[ResponseTool]>,
|
||||
include_mcp: bool,
|
||||
) -> Vec<Tool> {
|
||||
@@ -124,7 +124,7 @@ pub fn extract_tools_from_response_tools(
|
||||
///
|
||||
/// Common helper function to avoid duplication across sync and streaming paths
|
||||
/// in both harmony and regular responses implementations.
|
||||
pub async fn persist_response_if_needed(
|
||||
pub(crate) async fn persist_response_if_needed(
|
||||
conversation_storage: Arc<dyn ConversationStorage>,
|
||||
conversation_item_storage: Arc<dyn ConversationItemStorage>,
|
||||
response_storage: Arc<dyn ResponseStorage>,
|
||||
|
||||
@@ -14,7 +14,7 @@ use crate::routers::{
|
||||
};
|
||||
|
||||
/// Client acquisition stage: Get gRPC clients from selected workers
|
||||
pub struct ClientAcquisitionStage;
|
||||
pub(crate) struct ClientAcquisitionStage;
|
||||
|
||||
#[async_trait]
|
||||
impl PipelineStage for ClientAcquisitionStage {
|
||||
|
||||
@@ -16,7 +16,7 @@ use crate::{
|
||||
};
|
||||
|
||||
/// Dispatch metadata stage: Prepare metadata for dispatch
|
||||
pub struct DispatchMetadataStage;
|
||||
pub(crate) struct DispatchMetadataStage;
|
||||
|
||||
#[async_trait]
|
||||
impl PipelineStage for DispatchMetadataStage {
|
||||
|
||||
@@ -14,7 +14,7 @@ use crate::{
|
||||
///
|
||||
/// Used by both chat and generate request building stages when in PD mode.
|
||||
/// Only SGLang supports PD (prefill/decode) disaggregated mode.
|
||||
pub fn inject_bootstrap_metadata(
|
||||
pub(crate) fn inject_bootstrap_metadata(
|
||||
request: &mut ProtoGenerateRequest,
|
||||
prefill_worker: &Arc<dyn Worker>,
|
||||
) {
|
||||
|
||||
@@ -28,12 +28,12 @@ pub trait PipelineStage: Send + Sync {
|
||||
|
||||
mod client_acquisition;
|
||||
mod dispatch_metadata;
|
||||
pub mod helpers;
|
||||
pub(crate) mod helpers;
|
||||
mod request_execution;
|
||||
mod worker_selection;
|
||||
|
||||
// Export stage implementations
|
||||
pub use client_acquisition::ClientAcquisitionStage;
|
||||
pub use dispatch_metadata::DispatchMetadataStage;
|
||||
pub use request_execution::{ExecutionMode, RequestExecutionStage};
|
||||
pub use worker_selection::{WorkerSelectionMode, WorkerSelectionStage};
|
||||
pub(crate) use client_acquisition::ClientAcquisitionStage;
|
||||
pub(crate) use dispatch_metadata::DispatchMetadataStage;
|
||||
pub(crate) use request_execution::{ExecutionMode, RequestExecutionStage};
|
||||
pub(crate) use worker_selection::{WorkerSelectionMode, WorkerSelectionStage};
|
||||
|
||||
@@ -19,12 +19,12 @@ use crate::routers::{
|
||||
type StreamResult = Result<ProtoStream, Box<dyn std::error::Error + Send + Sync>>;
|
||||
|
||||
/// Request execution stage: Execute gRPC requests (single or dual dispatch)
|
||||
pub struct RequestExecutionStage {
|
||||
pub(crate) struct RequestExecutionStage {
|
||||
mode: ExecutionMode,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub enum ExecutionMode {
|
||||
pub(crate) enum ExecutionMode {
|
||||
/// Regular mode: single worker execution
|
||||
Single,
|
||||
/// PD mode: dual dispatch to prefill + decode workers
|
||||
|
||||
@@ -18,13 +18,13 @@ use crate::{
|
||||
};
|
||||
|
||||
/// Worker selection stage: Select appropriate worker(s) based on routing mode
|
||||
pub struct WorkerSelectionStage {
|
||||
pub(crate) struct WorkerSelectionStage {
|
||||
worker_registry: Arc<WorkerRegistry>,
|
||||
policy_registry: Arc<PolicyRegistry>,
|
||||
mode: WorkerSelectionMode,
|
||||
}
|
||||
|
||||
pub enum WorkerSelectionMode {
|
||||
pub(crate) enum WorkerSelectionMode {
|
||||
/// Regular mode: select single worker
|
||||
Regular,
|
||||
/// PD mode: select prefill + decode workers
|
||||
|
||||
@@ -4,14 +4,13 @@
|
||||
//! eliminating deep parameter passing chains and providing a single source of truth
|
||||
//! for request state.
|
||||
|
||||
use std::{collections::HashMap, sync::Arc};
|
||||
use std::sync::Arc;
|
||||
|
||||
use axum::http::HeaderMap;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
client::GrpcClient,
|
||||
proto_wrapper::{ProtoEmbedComplete, ProtoGenerateComplete, ProtoRequest, ProtoStream},
|
||||
proto_wrapper::{ProtoEmbedComplete, ProtoRequest, ProtoStream},
|
||||
};
|
||||
use crate::{
|
||||
core::{attach_guards_to_response, Worker, WorkerLoadGuard},
|
||||
@@ -32,14 +31,14 @@ use crate::{
|
||||
/// This is the single source of truth for all request state as it flows
|
||||
/// through the pipeline stages. Uses Rust's type system to enforce proper
|
||||
/// stage ordering at compile time.
|
||||
pub struct RequestContext {
|
||||
pub(crate) struct RequestContext {
|
||||
pub input: RequestInput,
|
||||
pub components: Arc<SharedComponents>,
|
||||
pub state: ProcessingState,
|
||||
}
|
||||
|
||||
/// Immutable request input
|
||||
pub struct RequestInput {
|
||||
pub(crate) struct RequestInput {
|
||||
pub request_type: RequestType,
|
||||
pub headers: Option<HeaderMap>,
|
||||
pub model_id: Option<String>,
|
||||
@@ -47,7 +46,7 @@ pub struct RequestInput {
|
||||
|
||||
/// Request type variants
|
||||
/// Using Arc instead of Box to enable cheap cloning for background tasks
|
||||
pub enum RequestType {
|
||||
pub(crate) enum RequestType {
|
||||
Chat(Arc<ChatCompletionRequest>),
|
||||
Generate(Arc<GenerateRequest>),
|
||||
Responses(Arc<ResponsesRequest>),
|
||||
@@ -56,15 +55,17 @@ pub enum RequestType {
|
||||
}
|
||||
|
||||
/// Shared components (injected once at creation)
|
||||
pub struct SharedComponents {
|
||||
pub(crate) struct SharedComponents {
|
||||
pub tokenizer_registry: Arc<TokenizerRegistry>,
|
||||
#[allow(dead_code)]
|
||||
pub tool_parser_factory: ToolParserFactory,
|
||||
#[allow(dead_code)]
|
||||
pub reasoning_parser_factory: ReasoningParserFactory,
|
||||
}
|
||||
|
||||
/// Mutable processing state (evolves through pipeline stages)
|
||||
#[derive(Default)]
|
||||
pub struct ProcessingState {
|
||||
pub(crate) struct ProcessingState {
|
||||
// Stage 1: Preparation outputs
|
||||
pub preparation: Option<PreparationOutput>,
|
||||
|
||||
@@ -92,7 +93,7 @@ pub struct ProcessingState {
|
||||
}
|
||||
|
||||
/// Output from preparation stage (Step 1)
|
||||
pub struct PreparationOutput {
|
||||
pub(crate) struct PreparationOutput {
|
||||
/// Original text (for chat) or resolved text (for generate)
|
||||
pub original_text: Option<String>,
|
||||
|
||||
@@ -116,6 +117,7 @@ pub struct PreparationOutput {
|
||||
pub selection_text: Option<String>,
|
||||
|
||||
/// Harmony messages for history tracking (Harmony only)
|
||||
#[allow(dead_code)]
|
||||
pub harmony_messages: Option<Vec<super::harmony::HarmonyMessage>>,
|
||||
|
||||
/// Stop token IDs for Harmony models
|
||||
@@ -123,7 +125,7 @@ pub struct PreparationOutput {
|
||||
}
|
||||
|
||||
/// Worker selection (Step 2)
|
||||
pub enum WorkerSelection {
|
||||
pub(crate) enum WorkerSelection {
|
||||
Single {
|
||||
worker: Arc<dyn Worker>,
|
||||
},
|
||||
@@ -134,7 +136,7 @@ pub enum WorkerSelection {
|
||||
}
|
||||
|
||||
/// Client selection (Step 3)
|
||||
pub enum ClientSelection {
|
||||
pub(crate) enum ClientSelection {
|
||||
Single {
|
||||
client: GrpcClient,
|
||||
},
|
||||
@@ -146,17 +148,18 @@ pub enum ClientSelection {
|
||||
|
||||
/// Dispatch metadata (Step 5)
|
||||
#[derive(Clone)]
|
||||
pub struct DispatchMetadata {
|
||||
pub(crate) struct DispatchMetadata {
|
||||
pub request_id: String,
|
||||
pub model: String,
|
||||
pub created: u64,
|
||||
pub weight_version: Option<String>,
|
||||
#[allow(dead_code)]
|
||||
pub is_streaming: bool,
|
||||
}
|
||||
|
||||
/// Load guards for worker load tracking
|
||||
/// Automatically decrements load when dropped
|
||||
pub enum LoadGuards {
|
||||
pub(crate) enum LoadGuards {
|
||||
Single(WorkerLoadGuard),
|
||||
Dual {
|
||||
prefill: WorkerLoadGuard,
|
||||
@@ -200,19 +203,10 @@ impl LoadGuards {
|
||||
|
||||
/// Response processing state (Step 6)
|
||||
#[derive(Default)]
|
||||
pub struct ResponseState {
|
||||
pub(crate) struct ResponseState {
|
||||
/// Stop sequence decoder
|
||||
pub stop_decoder: Option<StopSequenceDecoder>,
|
||||
|
||||
/// Per-index streaming state (for n>1 support)
|
||||
pub streaming: StreamingState,
|
||||
|
||||
/// Collected responses (non-streaming)
|
||||
pub collected: Option<Vec<ProtoGenerateComplete>>,
|
||||
|
||||
/// Collected embeddings (non-streaming)
|
||||
pub collected_embeddings: Option<Vec<ProtoEmbedComplete>>,
|
||||
|
||||
/// Execution result (streams from workers)
|
||||
pub execution_result: Option<ExecutionResult>,
|
||||
|
||||
@@ -221,32 +215,6 @@ pub struct ResponseState {
|
||||
|
||||
/// Responses API iteration result (Harmony only, for tool loop orchestration)
|
||||
pub responses_iteration_result: Option<super::harmony::ResponsesIterationResult>,
|
||||
|
||||
// Harmony-specific parser state
|
||||
/// Harmony parser for non-streaming (single parser for all indices)
|
||||
pub harmony_parser: Option<super::harmony::HarmonyParserAdapter>,
|
||||
|
||||
/// Harmony parsers for streaming (one per index for n>1 support)
|
||||
pub harmony_parser_per_index: Option<HashMap<usize, super::harmony::HarmonyParserAdapter>>,
|
||||
}
|
||||
|
||||
/// Streaming state (per-choice tracking)
|
||||
#[derive(Default)]
|
||||
pub struct StreamingState {
|
||||
pub is_firsts: HashMap<u32, bool>,
|
||||
pub stream_buffers: HashMap<u32, String>,
|
||||
pub finish_reasons: HashMap<u32, String>,
|
||||
pub matched_stops: HashMap<u32, Option<Value>>,
|
||||
pub prompt_tokens: HashMap<u32, u32>,
|
||||
pub completion_tokens: HashMap<u32, u32>,
|
||||
pub cached_tokens: HashMap<u32, u32>,
|
||||
|
||||
// Parser state (lazy initialization per index)
|
||||
pub reasoning_parsers:
|
||||
HashMap<u32, Arc<std::sync::Mutex<Box<dyn crate::reasoning_parser::ReasoningParser>>>>,
|
||||
pub tool_parsers:
|
||||
HashMap<u32, Arc<tokio::sync::Mutex<Box<dyn crate::tool_parser::ToolParser>>>>,
|
||||
pub has_tool_calls: HashMap<u32, bool>,
|
||||
}
|
||||
|
||||
impl RequestContext {
|
||||
@@ -340,11 +308,6 @@ impl RequestContext {
|
||||
}
|
||||
}
|
||||
|
||||
/// Get reference to original request (type-safe)
|
||||
pub fn request(&self) -> &RequestType {
|
||||
&self.input.request_type
|
||||
}
|
||||
|
||||
/// Get chat request (panics if not chat)
|
||||
pub fn chat_request(&self) -> &ChatCompletionRequest {
|
||||
match &self.input.request_type {
|
||||
@@ -377,14 +340,6 @@ impl RequestContext {
|
||||
}
|
||||
}
|
||||
|
||||
/// Get responses request (panics if not responses)
|
||||
pub fn responses_request(&self) -> &ResponsesRequest {
|
||||
match &self.input.request_type {
|
||||
RequestType::Responses(req) => req.as_ref(),
|
||||
_ => panic!("Expected responses request"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Get Arc clone of responses request (panics if not responses)
|
||||
pub fn responses_request_arc(&self) -> Arc<ResponsesRequest> {
|
||||
match &self.input.request_type {
|
||||
@@ -393,38 +348,6 @@ impl RequestContext {
|
||||
}
|
||||
}
|
||||
|
||||
/// Get embedding request (panics if not embedding)
|
||||
pub fn embedding_request(&self) -> &EmbeddingRequest {
|
||||
match &self.input.request_type {
|
||||
RequestType::Embedding(req) => req.as_ref(),
|
||||
_ => panic!("Expected embedding request"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Get Arc clone of embedding request (panics if not embedding)
|
||||
pub fn embedding_request_arc(&self) -> Arc<EmbeddingRequest> {
|
||||
match &self.input.request_type {
|
||||
RequestType::Embedding(req) => Arc::clone(req),
|
||||
_ => panic!("Expected embedding request"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Get classify request (panics if not classify)
|
||||
pub fn classify_request(&self) -> &ClassifyRequest {
|
||||
match &self.input.request_type {
|
||||
RequestType::Classify(req) => req.as_ref(),
|
||||
_ => panic!("Expected classify request"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Get Arc clone of classify request (panics if not classify)
|
||||
pub fn classify_request_arc(&self) -> Arc<ClassifyRequest> {
|
||||
match &self.input.request_type {
|
||||
RequestType::Classify(req) => Arc::clone(req),
|
||||
_ => panic!("Expected classify request"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if request is streaming
|
||||
pub fn is_streaming(&self) -> bool {
|
||||
match &self.input.request_type {
|
||||
@@ -446,10 +369,12 @@ impl RequestContext {
|
||||
}
|
||||
|
||||
impl WorkerSelection {
|
||||
#[allow(dead_code)]
|
||||
pub fn is_dual(&self) -> bool {
|
||||
matches!(self, Self::Dual { .. })
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn single(&self) -> Option<&Arc<dyn Worker>> {
|
||||
match self {
|
||||
Self::Single { worker } => Some(worker),
|
||||
@@ -476,6 +401,7 @@ impl WorkerSelection {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
#[allow(clippy::type_complexity)]
|
||||
pub fn dual(&self) -> Option<(&Arc<dyn Worker>, &Arc<dyn Worker>)> {
|
||||
match self {
|
||||
@@ -484,6 +410,7 @@ impl WorkerSelection {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn prefill_worker(&self) -> Option<&Arc<dyn Worker>> {
|
||||
match self {
|
||||
Self::Dual { prefill, .. } => Some(prefill),
|
||||
@@ -491,6 +418,7 @@ impl WorkerSelection {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn decode_worker(&self) -> Option<&Arc<dyn Worker>> {
|
||||
match self {
|
||||
Self::Dual { decode, .. } => Some(decode),
|
||||
@@ -500,6 +428,7 @@ impl WorkerSelection {
|
||||
}
|
||||
|
||||
impl ClientSelection {
|
||||
#[allow(dead_code)]
|
||||
pub fn is_dual(&self) -> bool {
|
||||
matches!(self, Self::Dual { .. })
|
||||
}
|
||||
@@ -518,6 +447,7 @@ impl ClientSelection {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn dual(&self) -> Option<(&GrpcClient, &GrpcClient)> {
|
||||
match self {
|
||||
Self::Dual { prefill, decode } => Some((prefill, decode)),
|
||||
@@ -532,6 +462,7 @@ impl ClientSelection {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn prefill_client(&self) -> Option<&GrpcClient> {
|
||||
match self {
|
||||
Self::Dual { prefill, .. } => Some(prefill),
|
||||
@@ -539,6 +470,7 @@ impl ClientSelection {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn prefill_client_mut(&mut self) -> Option<&mut GrpcClient> {
|
||||
match self {
|
||||
Self::Dual { prefill, .. } => Some(prefill),
|
||||
@@ -546,6 +478,7 @@ impl ClientSelection {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn decode_client(&self) -> Option<&GrpcClient> {
|
||||
match self {
|
||||
Self::Dual { decode, .. } => Some(decode),
|
||||
@@ -553,6 +486,7 @@ impl ClientSelection {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn decode_client_mut(&mut self) -> Option<&mut GrpcClient> {
|
||||
match self {
|
||||
Self::Dual { decode, .. } => Some(decode),
|
||||
@@ -563,7 +497,7 @@ impl ClientSelection {
|
||||
|
||||
/// Result of request execution (streams from workers)
|
||||
/// Uses ProtoStream to automatically abort on cancellation
|
||||
pub enum ExecutionResult {
|
||||
pub(crate) enum ExecutionResult {
|
||||
Single {
|
||||
stream: ProtoStream,
|
||||
},
|
||||
@@ -579,7 +513,7 @@ pub enum ExecutionResult {
|
||||
|
||||
/// Final processed response
|
||||
#[derive(Debug)]
|
||||
pub enum FinalResponse {
|
||||
pub(crate) enum FinalResponse {
|
||||
Chat(ChatCompletionResponse),
|
||||
/// Generate response is a Vec of GenerateResponse (n=1 returns single item, n>1 returns multiple)
|
||||
Generate(Vec<GenerateResponse>),
|
||||
|
||||
@@ -113,7 +113,7 @@ fn has_custom_tools(tool_types: &[&str]) -> bool {
|
||||
///
|
||||
/// Converts OpenAI-format requests into Harmony-encoded format with input_ids,
|
||||
/// stop tokens, and selection text for worker routing.
|
||||
pub struct HarmonyBuilder {
|
||||
pub(crate) struct HarmonyBuilder {
|
||||
encoding: &'static HarmonyEncoding,
|
||||
}
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ use crate::core::{Worker, WorkerRegistry};
|
||||
/// Harmony model detector
|
||||
///
|
||||
/// Detects if a model name indicates support for Harmony encoding/parsing.
|
||||
pub struct HarmonyDetector;
|
||||
pub(crate) struct HarmonyDetector;
|
||||
|
||||
impl HarmonyDetector {
|
||||
/// Check if a worker is a Harmony/GPT-OSS model.
|
||||
|
||||
@@ -29,28 +29,22 @@
|
||||
//! }
|
||||
//! ```
|
||||
|
||||
pub mod builder;
|
||||
pub mod detector;
|
||||
pub mod parser;
|
||||
pub mod processor;
|
||||
pub mod responses;
|
||||
pub mod stages;
|
||||
pub mod streaming;
|
||||
pub mod types;
|
||||
pub(crate) mod builder;
|
||||
pub(crate) mod detector;
|
||||
pub(crate) mod parser;
|
||||
pub(crate) mod processor;
|
||||
pub(crate) mod responses;
|
||||
pub(crate) mod stages;
|
||||
pub(crate) mod streaming;
|
||||
pub(crate) mod types;
|
||||
|
||||
// Re-export main types for convenience
|
||||
pub use builder::HarmonyBuilder;
|
||||
pub use detector::HarmonyDetector;
|
||||
pub use parser::HarmonyParserAdapter;
|
||||
pub use processor::{HarmonyResponseProcessor, ResponsesIterationResult};
|
||||
pub use responses::{
|
||||
// Re-export types that are accessed via harmony::TypeName
|
||||
pub(crate) use builder::HarmonyBuilder;
|
||||
pub(crate) use detector::HarmonyDetector;
|
||||
pub(crate) use parser::HarmonyParserAdapter;
|
||||
pub(crate) use processor::{HarmonyResponseProcessor, ResponsesIterationResult};
|
||||
pub(crate) use responses::{
|
||||
serve_harmony_responses, serve_harmony_responses_stream, HarmonyResponsesContext,
|
||||
};
|
||||
pub use stages::{
|
||||
HarmonyPreparationStage, HarmonyRequestBuildingStage, HarmonyResponseProcessingStage,
|
||||
};
|
||||
pub use streaming::HarmonyStreamingProcessor;
|
||||
pub use types::{
|
||||
FunctionDelta, HarmonyBuildOutput, HarmonyChannelDelta, HarmonyChannelOutput, HarmonyMessage,
|
||||
ToolCallDelta,
|
||||
};
|
||||
pub(crate) use streaming::HarmonyStreamingProcessor;
|
||||
pub(crate) use types::HarmonyMessage;
|
||||
|
||||
@@ -20,7 +20,7 @@ fn get_harmony_encoding() -> &'static HarmonyEncoding {
|
||||
///
|
||||
/// Wraps openai_harmony::StreamableParser and provides methods for parsing
|
||||
/// complete responses and streaming chunks.
|
||||
pub struct HarmonyParserAdapter {
|
||||
pub(crate) struct HarmonyParserAdapter {
|
||||
parser: StreamableParser,
|
||||
prev_recipient: Option<String>,
|
||||
reasoning_token_count: u32,
|
||||
@@ -517,6 +517,7 @@ impl HarmonyParserAdapter {
|
||||
/// Reset parser state
|
||||
///
|
||||
/// Resets the parser to initial state for reuse
|
||||
#[allow(dead_code)]
|
||||
pub fn reset(&mut self) -> Result<(), String> {
|
||||
// Create a new parser instance (StreamableParser doesn't have a reset method)
|
||||
let encoding = get_harmony_encoding();
|
||||
|
||||
@@ -29,7 +29,7 @@ use crate::{
|
||||
///
|
||||
/// Collects all output tokens from execution and parses them using
|
||||
/// HarmonyParserAdapter to extract the complete response.
|
||||
pub struct HarmonyResponseProcessor;
|
||||
pub(crate) struct HarmonyResponseProcessor;
|
||||
|
||||
impl HarmonyResponseProcessor {
|
||||
/// Create a new Harmony response processor
|
||||
@@ -155,7 +155,7 @@ impl Default for HarmonyResponseProcessor {
|
||||
///
|
||||
/// Used by the MCP tool loop to determine whether to continue
|
||||
/// executing tools or return the final response.
|
||||
pub enum ResponsesIterationResult {
|
||||
pub(crate) enum ResponsesIterationResult {
|
||||
/// Tool calls found in commentary channel - continue MCP loop
|
||||
ToolCallsFound {
|
||||
tool_calls: Vec<ToolCall>,
|
||||
|
||||
@@ -2,8 +2,6 @@
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use crate::{
|
||||
data_connector::{ConversationItemStorage, ConversationStorage, ResponseStorage},
|
||||
mcp::McpManager,
|
||||
@@ -15,7 +13,7 @@ use crate::{
|
||||
/// Contains all dependencies needed for multi-turn Responses API execution.
|
||||
/// Cheap to clone (all Arc references).
|
||||
#[derive(Clone)]
|
||||
pub struct HarmonyResponsesContext {
|
||||
pub(crate) struct HarmonyResponsesContext {
|
||||
/// Pipeline for executing Harmony requests
|
||||
pub pipeline: Arc<RequestPipeline>,
|
||||
|
||||
@@ -33,9 +31,6 @@ pub struct HarmonyResponsesContext {
|
||||
|
||||
/// Conversation item storage for persisting conversation items
|
||||
pub conversation_item_storage: Arc<dyn ConversationItemStorage>,
|
||||
|
||||
/// Optional streaming sender (for future streaming support)
|
||||
pub stream_tx: Option<mpsc::UnboundedSender<Result<String, String>>>,
|
||||
}
|
||||
|
||||
impl HarmonyResponsesContext {
|
||||
@@ -55,28 +50,6 @@ impl HarmonyResponsesContext {
|
||||
response_storage,
|
||||
conversation_storage,
|
||||
conversation_item_storage,
|
||||
stream_tx: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Create with streaming support
|
||||
pub fn with_streaming(
|
||||
pipeline: Arc<RequestPipeline>,
|
||||
components: Arc<SharedComponents>,
|
||||
mcp_manager: Arc<McpManager>,
|
||||
response_storage: Arc<dyn ResponseStorage>,
|
||||
conversation_storage: Arc<dyn ConversationStorage>,
|
||||
conversation_item_storage: Arc<dyn ConversationItemStorage>,
|
||||
stream_tx: mpsc::UnboundedSender<Result<String, String>>,
|
||||
) -> Self {
|
||||
Self {
|
||||
pipeline,
|
||||
components,
|
||||
mcp_manager,
|
||||
response_storage,
|
||||
conversation_storage,
|
||||
conversation_item_storage,
|
||||
stream_tx: Some(stream_tx),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -20,7 +20,7 @@ use crate::{
|
||||
/// Tool execution result
|
||||
///
|
||||
/// Contains the result of executing a single MCP tool.
|
||||
pub struct ToolResult {
|
||||
pub(crate) struct ToolResult {
|
||||
/// Tool call ID (for matching with request)
|
||||
pub call_id: String,
|
||||
|
||||
@@ -202,7 +202,7 @@ pub(super) async fn execute_mcp_tools(
|
||||
///
|
||||
/// Converts MCP Tool entries (from rmcp SDK) to ResponseTool format so the model
|
||||
/// knows about available MCP tools when making tool calls.
|
||||
pub fn convert_mcp_tools_to_response_tools(mcp_tools: &[mcp::Tool]) -> Vec<ResponseTool> {
|
||||
pub(crate) fn convert_mcp_tools_to_response_tools(mcp_tools: &[mcp::Tool]) -> Vec<ResponseTool> {
|
||||
mcp_tools
|
||||
.iter()
|
||||
.map(|tool_info| ResponseTool {
|
||||
|
||||
@@ -19,14 +19,14 @@
|
||||
//! - `execution` - MCP tool execution logic
|
||||
//! - `common` - Shared helpers and state tracking
|
||||
|
||||
mod common;
|
||||
mod context;
|
||||
mod execution;
|
||||
mod non_streaming;
|
||||
mod streaming;
|
||||
pub(crate) mod common;
|
||||
pub(crate) mod context;
|
||||
pub(crate) mod execution;
|
||||
pub(crate) mod non_streaming;
|
||||
pub(crate) mod streaming;
|
||||
|
||||
// Public exports
|
||||
pub use context::HarmonyResponsesContext;
|
||||
pub use execution::{convert_mcp_tools_to_response_tools, ToolResult};
|
||||
pub use non_streaming::serve_harmony_responses;
|
||||
pub use streaming::serve_harmony_responses_stream;
|
||||
// Re-export types accessed via harmony::responses::TypeName
|
||||
pub(crate) use context::HarmonyResponsesContext;
|
||||
pub(crate) use execution::ToolResult;
|
||||
pub(crate) use non_streaming::serve_harmony_responses;
|
||||
pub(crate) use streaming::serve_harmony_responses_stream;
|
||||
|
||||
@@ -46,7 +46,7 @@ use crate::{
|
||||
/// - Build next request with tool results
|
||||
/// - Repeat from step 1 (full pipeline re-execution)
|
||||
/// 4. If no tool calls, return final response
|
||||
pub async fn serve_harmony_responses(
|
||||
pub(crate) async fn serve_harmony_responses(
|
||||
ctx: &HarmonyResponsesContext,
|
||||
request: ResponsesRequest,
|
||||
) -> Result<ResponsesResponse, Response> {
|
||||
|
||||
@@ -36,7 +36,7 @@ use crate::{
|
||||
///
|
||||
/// This is the streaming equivalent of `serve_harmony_responses()`.
|
||||
/// Emits SSE events for lifecycle, MCP list_tools, and per-iteration streaming.
|
||||
pub async fn serve_harmony_responses_stream(
|
||||
pub(crate) async fn serve_harmony_responses_stream(
|
||||
ctx: &HarmonyResponsesContext,
|
||||
request: ResponsesRequest,
|
||||
) -> Response {
|
||||
|
||||
@@ -5,10 +5,10 @@
|
||||
//! - HarmonyRequestBuildingStage: Token-based request building
|
||||
//! - HarmonyResponseProcessingStage: Harmony channel parsing
|
||||
|
||||
pub mod preparation;
|
||||
pub mod request_building;
|
||||
pub mod response_processing;
|
||||
pub(crate) mod preparation;
|
||||
pub(crate) mod request_building;
|
||||
pub(crate) mod response_processing;
|
||||
|
||||
pub use preparation::HarmonyPreparationStage;
|
||||
pub use request_building::HarmonyRequestBuildingStage;
|
||||
pub use response_processing::HarmonyResponseProcessingStage;
|
||||
pub(crate) use preparation::HarmonyPreparationStage;
|
||||
pub(crate) use request_building::HarmonyRequestBuildingStage;
|
||||
pub(crate) use response_processing::HarmonyResponseProcessingStage;
|
||||
|
||||
@@ -26,7 +26,7 @@ use crate::{
|
||||
///
|
||||
/// Replaces the regular PreparationStage for Harmony models.
|
||||
/// Converts chat/generate requests to Harmony-encoded token_ids and extraction_text.
|
||||
pub struct HarmonyPreparationStage {
|
||||
pub(crate) struct HarmonyPreparationStage {
|
||||
builder: HarmonyBuilder,
|
||||
}
|
||||
|
||||
@@ -387,7 +387,9 @@ impl HarmonyPreparationStage {
|
||||
/// - Without reasoning: triggers on `<|channel|>final` (goes directly to final channel)
|
||||
///
|
||||
/// This is used for the Responses API text.format field (json_object or json_schema).
|
||||
pub fn build_text_format_structural_tag(schema: &serde_json::Value) -> Result<String, String> {
|
||||
pub(crate) fn build_text_format_structural_tag(
|
||||
schema: &serde_json::Value,
|
||||
) -> Result<String, String> {
|
||||
let structural_tag = json!({
|
||||
"format": {
|
||||
"type": "triggered_tags",
|
||||
|
||||
@@ -18,7 +18,7 @@ use crate::routers::{
|
||||
///
|
||||
/// Takes the Harmony-encoded input_ids from preparation and builds a proto::GenerateRequest.
|
||||
/// Unlike regular request building, this uses token_ids directly (Harmony encoding handles messages).
|
||||
pub struct HarmonyRequestBuildingStage {
|
||||
pub(crate) struct HarmonyRequestBuildingStage {
|
||||
inject_pd_metadata: bool,
|
||||
}
|
||||
|
||||
|
||||
@@ -19,7 +19,7 @@ use crate::routers::{
|
||||
///
|
||||
/// Takes output tokens from execution and parses them using HarmonyParserAdapter
|
||||
/// to extract analysis, tool calls, and final response text from Harmony channels.
|
||||
pub struct HarmonyResponseProcessingStage {
|
||||
pub(crate) struct HarmonyResponseProcessingStage {
|
||||
processor: HarmonyResponseProcessor,
|
||||
streaming_processor: Arc<HarmonyStreamingProcessor>,
|
||||
}
|
||||
|
||||
@@ -105,7 +105,7 @@ impl ToolCallMode {
|
||||
///
|
||||
/// Returns an SSE stream that parses Harmony tokens incrementally and
|
||||
/// emits ChatCompletionChunk events for streaming responses.
|
||||
pub struct HarmonyStreamingProcessor;
|
||||
pub(crate) struct HarmonyStreamingProcessor;
|
||||
|
||||
impl HarmonyStreamingProcessor {
|
||||
/// Create a new Harmony streaming processor
|
||||
|
||||
@@ -10,11 +10,12 @@ use crate::protocols::common::ToolCall;
|
||||
///
|
||||
/// Represents messages in the Harmony encoding format with role and content.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct HarmonyMessage {
|
||||
pub(crate) struct HarmonyMessage {
|
||||
pub role: String,
|
||||
pub content: String,
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
impl HarmonyMessage {
|
||||
pub fn new(role: impl Into<String>, content: impl Into<String>) -> Self {
|
||||
Self {
|
||||
@@ -67,7 +68,7 @@ impl HarmonyMessage {
|
||||
/// Contains the encoded input_ids, stop tokens, selection text for worker routing,
|
||||
/// and the Harmony message history.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct HarmonyBuildOutput {
|
||||
pub(crate) struct HarmonyBuildOutput {
|
||||
/// Encoded token IDs to send to the model
|
||||
pub input_ids: Vec<u32>,
|
||||
|
||||
@@ -85,7 +86,7 @@ pub struct HarmonyBuildOutput {
|
||||
///
|
||||
/// Represents the complete response after parsing analysis, commentary, and final channels.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct HarmonyChannelOutput {
|
||||
pub(crate) struct HarmonyChannelOutput {
|
||||
/// Analysis/reasoning content (from analysis channel)
|
||||
pub analysis: Option<String>,
|
||||
|
||||
@@ -109,7 +110,8 @@ pub struct HarmonyChannelOutput {
|
||||
///
|
||||
/// Represents incremental updates as tokens are parsed from the stream.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct HarmonyChannelDelta {
|
||||
#[allow(dead_code)]
|
||||
pub(crate) struct HarmonyChannelDelta {
|
||||
/// Delta for analysis/reasoning content
|
||||
pub analysis_delta: Option<String>,
|
||||
|
||||
@@ -125,7 +127,7 @@ pub struct HarmonyChannelDelta {
|
||||
|
||||
/// Tool call delta for streaming
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ToolCallDelta {
|
||||
pub(crate) struct ToolCallDelta {
|
||||
pub index: usize,
|
||||
pub id: Option<String>,
|
||||
pub function: Option<FunctionDelta>,
|
||||
@@ -133,7 +135,7 @@ pub struct ToolCallDelta {
|
||||
|
||||
/// Function call delta for streaming
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct FunctionDelta {
|
||||
pub(crate) struct FunctionDelta {
|
||||
pub name: Option<String>,
|
||||
pub arguments: Option<String>,
|
||||
}
|
||||
|
||||
@@ -2,21 +2,22 @@
|
||||
|
||||
use crate::{grpc_client::sglang_proto::MultimodalInputs, protocols::common::StringOrArray};
|
||||
|
||||
pub mod client;
|
||||
pub mod common;
|
||||
pub mod context;
|
||||
pub mod harmony;
|
||||
pub mod pd_router;
|
||||
pub mod pipeline;
|
||||
pub mod proto_wrapper;
|
||||
pub mod regular;
|
||||
pub mod router;
|
||||
pub mod utils;
|
||||
pub mod client; // Used by core/
|
||||
pub(crate) mod common;
|
||||
pub(crate) mod context;
|
||||
pub(crate) mod harmony;
|
||||
pub(crate) mod pd_router; // Used by routers/factory
|
||||
pub(crate) mod pipeline;
|
||||
pub(crate) mod proto_wrapper;
|
||||
pub(crate) mod regular;
|
||||
pub(crate) mod router; // Used by routers/factory
|
||||
pub(crate) mod utils; // Used by routers/http
|
||||
|
||||
/// Processed chat messages ready for gRPC generation
|
||||
#[derive(Debug)]
|
||||
pub struct ProcessedMessages {
|
||||
pub(crate) struct ProcessedMessages {
|
||||
pub text: String,
|
||||
pub multimodal_inputs: Option<MultimodalInputs>,
|
||||
#[allow(dead_code)]
|
||||
pub stop_sequences: Option<StringOrArray>,
|
||||
}
|
||||
|
||||
@@ -48,7 +48,7 @@ use crate::{
|
||||
/// Orchestrates all stages from request preparation to response delivery.
|
||||
/// Configured differently for regular vs PD mode.
|
||||
#[derive(Clone)]
|
||||
pub struct RequestPipeline {
|
||||
pub(crate) struct RequestPipeline {
|
||||
stages: Arc<Vec<Box<dyn PipelineStage>>>,
|
||||
/// Backend type for metrics labeling
|
||||
backend_type: &'static str,
|
||||
@@ -129,6 +129,7 @@ impl RequestPipeline {
|
||||
}
|
||||
|
||||
/// Create a Harmony PD (prefill-decode) pipeline
|
||||
#[allow(dead_code)]
|
||||
pub fn new_harmony_pd(
|
||||
worker_registry: Arc<WorkerRegistry>,
|
||||
policy_registry: Arc<PolicyRegistry>,
|
||||
@@ -369,9 +370,6 @@ impl RequestPipeline {
|
||||
components: Arc<SharedComponents>,
|
||||
) -> Response {
|
||||
let start = Instant::now();
|
||||
// Clone model_id for metrics before moving into context
|
||||
// GenerateRequest doesn't have a model field, so we use model_id
|
||||
let model_for_metrics = model_id.clone();
|
||||
let streaming = request.stream;
|
||||
|
||||
// Record request start
|
||||
@@ -379,12 +377,12 @@ impl RequestPipeline {
|
||||
metrics_labels::ROUTER_GRPC,
|
||||
self.backend_type,
|
||||
metrics_labels::CONNECTION_GRPC,
|
||||
model_for_metrics.as_deref().unwrap_or("unknown"),
|
||||
model_id.as_deref().unwrap_or(UNKNOWN_MODEL_ID),
|
||||
metrics_labels::ENDPOINT_GENERATE,
|
||||
bool_to_static_str(streaming),
|
||||
);
|
||||
|
||||
let mut ctx = RequestContext::for_generate(request, headers, model_id, components);
|
||||
let mut ctx = RequestContext::for_generate(request, headers, model_id.clone(), components);
|
||||
|
||||
for stage in self.stages.iter() {
|
||||
match stage.execute(&mut ctx).await {
|
||||
@@ -393,7 +391,7 @@ impl RequestPipeline {
|
||||
metrics_labels::ROUTER_GRPC,
|
||||
self.backend_type,
|
||||
metrics_labels::CONNECTION_GRPC,
|
||||
model_for_metrics.as_deref().unwrap_or("unknown"),
|
||||
model_id.as_deref().unwrap_or(UNKNOWN_MODEL_ID),
|
||||
metrics_labels::ENDPOINT_GENERATE,
|
||||
start.elapsed(),
|
||||
);
|
||||
@@ -405,7 +403,7 @@ impl RequestPipeline {
|
||||
metrics_labels::ROUTER_GRPC,
|
||||
self.backend_type,
|
||||
metrics_labels::CONNECTION_GRPC,
|
||||
model_for_metrics.as_deref().unwrap_or("unknown"),
|
||||
model_id.as_deref().unwrap_or(UNKNOWN_MODEL_ID),
|
||||
metrics_labels::ENDPOINT_GENERATE,
|
||||
error_type_from_status(response.status()),
|
||||
);
|
||||
@@ -425,7 +423,7 @@ impl RequestPipeline {
|
||||
metrics_labels::ROUTER_GRPC,
|
||||
self.backend_type,
|
||||
metrics_labels::CONNECTION_GRPC,
|
||||
model_for_metrics.as_deref().unwrap_or("unknown"),
|
||||
model_id.as_deref().unwrap_or(UNKNOWN_MODEL_ID),
|
||||
metrics_labels::ENDPOINT_GENERATE,
|
||||
start.elapsed(),
|
||||
);
|
||||
@@ -442,7 +440,7 @@ impl RequestPipeline {
|
||||
metrics_labels::ROUTER_GRPC,
|
||||
self.backend_type,
|
||||
metrics_labels::CONNECTION_GRPC,
|
||||
model_for_metrics.as_deref().unwrap_or("unknown"),
|
||||
model_id.as_deref().unwrap_or(UNKNOWN_MODEL_ID),
|
||||
metrics_labels::ENDPOINT_GENERATE,
|
||||
metrics_labels::ERROR_INTERNAL,
|
||||
);
|
||||
@@ -457,7 +455,7 @@ impl RequestPipeline {
|
||||
metrics_labels::ROUTER_GRPC,
|
||||
self.backend_type,
|
||||
metrics_labels::CONNECTION_GRPC,
|
||||
model_for_metrics.as_deref().unwrap_or("unknown"),
|
||||
model_id.as_deref().unwrap_or(UNKNOWN_MODEL_ID),
|
||||
metrics_labels::ENDPOINT_GENERATE,
|
||||
metrics_labels::ERROR_INTERNAL,
|
||||
);
|
||||
@@ -541,9 +539,7 @@ impl RequestPipeline {
|
||||
ctx.state.response.final_response
|
||||
);
|
||||
match ctx.state.response.final_response {
|
||||
Some(FinalResponse::Embedding(_)) => {
|
||||
error!("execute_embeddings: Embedding FinalResponse found, but pipeline finished without returning response directly. This should be handled by the last stage.");
|
||||
// Already handled in ResponseProcessingStage, but just in case
|
||||
Some(FinalResponse::Embedding(response)) => {
|
||||
Metrics::record_router_duration(
|
||||
metrics_labels::ROUTER_GRPC,
|
||||
self.backend_type,
|
||||
@@ -552,11 +548,7 @@ impl RequestPipeline {
|
||||
metrics_labels::ENDPOINT_EMBEDDINGS,
|
||||
start.elapsed(),
|
||||
);
|
||||
// The response should have been returned by the last stage
|
||||
error::internal_error(
|
||||
"pipeline_fallthrough",
|
||||
"Pipeline finished without returning response",
|
||||
)
|
||||
axum::Json(response).into_response()
|
||||
}
|
||||
Some(_) => {
|
||||
error!(function = "execute_embeddings", "Wrong response type");
|
||||
@@ -647,8 +639,7 @@ impl RequestPipeline {
|
||||
ctx.state.response.final_response
|
||||
);
|
||||
match ctx.state.response.final_response {
|
||||
Some(FinalResponse::Classify(_)) => {
|
||||
error!("execute_classify: Classify FinalResponse found, but pipeline finished without returning response directly. This should be handled by the last stage.");
|
||||
Some(FinalResponse::Classify(response)) => {
|
||||
Metrics::record_router_duration(
|
||||
metrics_labels::ROUTER_GRPC,
|
||||
self.backend_type,
|
||||
@@ -657,10 +648,7 @@ impl RequestPipeline {
|
||||
metrics_labels::ENDPOINT_CLASSIFY,
|
||||
start.elapsed(),
|
||||
);
|
||||
error::internal_error(
|
||||
"pipeline_fallthrough",
|
||||
"Pipeline finished without returning response",
|
||||
)
|
||||
axum::Json(response).into_response()
|
||||
}
|
||||
Some(_) => {
|
||||
error!(function = "execute_classify", "Wrong response type");
|
||||
|
||||
@@ -20,20 +20,7 @@ pub enum ProtoRequest {
|
||||
}
|
||||
|
||||
impl ProtoRequest {
|
||||
pub fn as_generate(&self) -> &ProtoGenerateRequest {
|
||||
match self {
|
||||
Self::Generate(req) => req,
|
||||
_ => panic!("Expected Generate request"),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn as_embed(&self) -> &ProtoEmbedRequest {
|
||||
match self {
|
||||
Self::Embed(req) => req,
|
||||
_ => panic!("Expected Embed request"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Get request ID from either variant
|
||||
pub fn request_id(&self) -> &str {
|
||||
match self {
|
||||
Self::Generate(req) => req.request_id(),
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
//! This module contains all code specific to regular tokenizer-based models,
|
||||
//! including pipeline stages, response processing, and streaming.
|
||||
|
||||
pub mod processor;
|
||||
pub mod responses;
|
||||
pub mod stages;
|
||||
pub mod streaming;
|
||||
pub(crate) mod processor;
|
||||
pub(crate) mod responses;
|
||||
pub(crate) mod stages;
|
||||
pub(crate) mod streaming;
|
||||
|
||||
@@ -34,7 +34,7 @@ use crate::{
|
||||
|
||||
/// Unified response processor for both routers
|
||||
#[derive(Clone)]
|
||||
pub struct ResponseProcessor {
|
||||
pub(crate) struct ResponseProcessor {
|
||||
pub tool_parser_factory: ToolParserFactory,
|
||||
pub reasoning_parser_factory: ReasoningParserFactory,
|
||||
pub configured_tool_parser: Option<String>,
|
||||
|
||||
@@ -19,7 +19,7 @@ use crate::{
|
||||
///
|
||||
/// This struct enables cancelling both the Rust task AND the Python scheduler processing.
|
||||
/// The client field is lazily initialized during pipeline execution.
|
||||
pub struct BackgroundTaskInfo {
|
||||
pub(crate) struct BackgroundTaskInfo {
|
||||
/// Tokio task handle for aborting the Rust task
|
||||
pub handle: JoinHandle<()>,
|
||||
/// gRPC request_id sent to Python scheduler (chatcmpl-* prefix)
|
||||
@@ -32,7 +32,7 @@ pub struct BackgroundTaskInfo {
|
||||
///
|
||||
/// All fields are Arc/shared references, so cloning this context is cheap.
|
||||
#[derive(Clone)]
|
||||
pub struct ResponsesContext {
|
||||
pub(crate) struct ResponsesContext {
|
||||
/// Chat pipeline for executing requests
|
||||
pub pipeline: Arc<RequestPipeline>,
|
||||
|
||||
@@ -40,6 +40,7 @@ pub struct ResponsesContext {
|
||||
pub components: Arc<SharedComponents>,
|
||||
|
||||
/// Worker registry for validation
|
||||
#[allow(dead_code)]
|
||||
pub worker_registry: Arc<WorkerRegistry>,
|
||||
|
||||
/// Response storage backend
|
||||
|
||||
@@ -33,7 +33,7 @@ use crate::{
|
||||
/// - `tools` → function tools extracted from ResponseTools
|
||||
/// - `tool_choice` → passed through from request
|
||||
/// - Response-specific fields (previous_response_id, conversation) are handled by router
|
||||
pub fn responses_to_chat(req: &ResponsesRequest) -> Result<ChatCompletionRequest, String> {
|
||||
pub(crate) fn responses_to_chat(req: &ResponsesRequest) -> Result<ChatCompletionRequest, String> {
|
||||
let mut messages = Vec::new();
|
||||
|
||||
// 1. Add system message if instructions provided
|
||||
@@ -271,7 +271,7 @@ fn map_text_to_response_format(text: &Option<TextConfig>) -> Option<ResponseForm
|
||||
/// - `choices[0].message` → `output` array (convert to ResponseOutputItem::Message)
|
||||
/// - `choices[0].finish_reason` → determines `status` (stop/length → Completed)
|
||||
/// - `created` timestamp → `created_at`
|
||||
pub fn chat_to_responses(
|
||||
pub(crate) fn chat_to_responses(
|
||||
chat_resp: &ChatCompletionResponse,
|
||||
original_req: &ResponsesRequest,
|
||||
response_id_override: Option<String>,
|
||||
|
||||
@@ -46,7 +46,7 @@ use crate::{
|
||||
/// Main handler for POST /v1/responses
|
||||
///
|
||||
/// Validates request, determines execution mode (sync/streaming), and delegates
|
||||
pub async fn route_responses(
|
||||
pub(crate) async fn route_responses(
|
||||
ctx: &ResponsesContext,
|
||||
request: Arc<ResponsesRequest>,
|
||||
headers: Option<http::HeaderMap>,
|
||||
|
||||
@@ -19,5 +19,5 @@ mod non_streaming;
|
||||
mod streaming;
|
||||
|
||||
// Public exports
|
||||
pub use context::{BackgroundTaskInfo, ResponsesContext};
|
||||
pub use handlers::route_responses;
|
||||
pub(crate) use context::ResponsesContext;
|
||||
pub(crate) use handlers::route_responses;
|
||||
|
||||
@@ -7,6 +7,6 @@ mod preparation;
|
||||
mod request_building;
|
||||
mod response_processing;
|
||||
|
||||
pub use preparation::ChatPreparationStage;
|
||||
pub use request_building::ChatRequestBuildingStage;
|
||||
pub use response_processing::ChatResponseProcessingStage;
|
||||
pub(crate) use preparation::ChatPreparationStage;
|
||||
pub(crate) use request_building::ChatRequestBuildingStage;
|
||||
pub(crate) use response_processing::ChatResponseProcessingStage;
|
||||
|
||||
@@ -22,7 +22,7 @@ use crate::{
|
||||
///
|
||||
/// Extracts chat-specific preparation logic from the old unified PreparationStage.
|
||||
/// This is a direct extraction without architectural changes.
|
||||
pub struct ChatPreparationStage;
|
||||
pub(crate) struct ChatPreparationStage;
|
||||
|
||||
#[async_trait]
|
||||
impl PipelineStage for ChatPreparationStage {
|
||||
|
||||
@@ -18,7 +18,7 @@ use crate::routers::{
|
||||
/// Chat request building stage
|
||||
///
|
||||
/// Extracts chat-specific request building logic from the old unified RequestBuildingStage.
|
||||
pub struct ChatRequestBuildingStage {
|
||||
pub(crate) struct ChatRequestBuildingStage {
|
||||
inject_pd_metadata: bool,
|
||||
}
|
||||
|
||||
|
||||
@@ -19,9 +19,7 @@ use crate::routers::{
|
||||
};
|
||||
|
||||
/// Chat response processing stage
|
||||
///
|
||||
/// Extracts chat-specific response processing logic from the old unified ResponseProcessingStage.
|
||||
pub struct ChatResponseProcessingStage {
|
||||
pub(crate) struct ChatResponseProcessingStage {
|
||||
processor: processor::ResponseProcessor,
|
||||
streaming_processor: Arc<streaming::StreamingProcessor>,
|
||||
}
|
||||
|
||||
@@ -4,6 +4,6 @@
|
||||
//! as the scheduler treats classify as an embedding request and returns logits.
|
||||
//! Only response processing is classify-specific (softmax + label mapping).
|
||||
|
||||
pub mod response_processing;
|
||||
pub(crate) mod response_processing;
|
||||
|
||||
pub use response_processing::ClassifyResponseProcessingStage;
|
||||
pub(crate) use response_processing::ClassifyResponseProcessingStage;
|
||||
|
||||
@@ -10,10 +10,7 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use axum::{
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use axum::response::Response;
|
||||
use tracing::error;
|
||||
|
||||
use crate::{
|
||||
@@ -37,7 +34,7 @@ use crate::{
|
||||
///
|
||||
/// The stage is stateless - id2label mapping is obtained from the
|
||||
/// selected worker's model card at runtime.
|
||||
pub struct ClassifyResponseProcessingStage;
|
||||
pub(crate) struct ClassifyResponseProcessingStage;
|
||||
|
||||
impl ClassifyResponseProcessingStage {
|
||||
/// Create a new classify response processing stage.
|
||||
@@ -205,11 +202,10 @@ impl PipelineStage for ClassifyResponseProcessingStage {
|
||||
usage,
|
||||
);
|
||||
|
||||
// Store in context
|
||||
ctx.state.response.final_response = Some(FinalResponse::Classify(response.clone()));
|
||||
// Store in context for pipeline to extract
|
||||
ctx.state.response.final_response = Some(FinalResponse::Classify(response));
|
||||
|
||||
// Return HTTP response
|
||||
Ok(Some(Json(response).into_response()))
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn name(&self) -> &'static str {
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
pub mod preparation;
|
||||
pub mod request_building;
|
||||
pub mod response_processing;
|
||||
pub(crate) mod preparation;
|
||||
pub(crate) mod request_building;
|
||||
pub(crate) mod response_processing;
|
||||
|
||||
@@ -16,7 +16,7 @@ use crate::{
|
||||
},
|
||||
};
|
||||
|
||||
pub struct EmbeddingPreparationStage;
|
||||
pub(crate) struct EmbeddingPreparationStage;
|
||||
|
||||
impl EmbeddingPreparationStage {
|
||||
pub fn new() -> Self {
|
||||
|
||||
@@ -15,7 +15,7 @@ use crate::routers::{
|
||||
};
|
||||
|
||||
/// Request building stage for embedding requests
|
||||
pub struct EmbeddingRequestBuildingStage;
|
||||
pub(crate) struct EmbeddingRequestBuildingStage;
|
||||
|
||||
impl EmbeddingRequestBuildingStage {
|
||||
pub fn new() -> Self {
|
||||
|
||||
@@ -1,10 +1,7 @@
|
||||
//! Response processing stage for embedding requests
|
||||
|
||||
use async_trait::async_trait;
|
||||
use axum::{
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use axum::response::Response;
|
||||
use tracing::error;
|
||||
|
||||
use crate::{
|
||||
@@ -20,7 +17,7 @@ use crate::{
|
||||
};
|
||||
|
||||
/// Response processing stage for embedding requests
|
||||
pub struct EmbeddingResponseProcessingStage;
|
||||
pub(crate) struct EmbeddingResponseProcessingStage;
|
||||
|
||||
impl EmbeddingResponseProcessingStage {
|
||||
pub fn new() -> Self {
|
||||
@@ -65,12 +62,10 @@ impl PipelineStage for EmbeddingResponseProcessingStage {
|
||||
.convert_response(ctx, proto_response)
|
||||
.map_err(|boxed_err| *boxed_err)?;
|
||||
|
||||
// Store in context
|
||||
ctx.state.response.final_response =
|
||||
Some(FinalResponse::Embedding(embedding_response.clone()));
|
||||
// Store in context for pipeline to extract
|
||||
ctx.state.response.final_response = Some(FinalResponse::Embedding(embedding_response));
|
||||
|
||||
// Return the HTTP response directly
|
||||
Ok(Some(Json(embedding_response).into_response()))
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn name(&self) -> &'static str {
|
||||
|
||||
@@ -7,6 +7,6 @@ mod preparation;
|
||||
mod request_building;
|
||||
mod response_processing;
|
||||
|
||||
pub use preparation::GeneratePreparationStage;
|
||||
pub use request_building::GenerateRequestBuildingStage;
|
||||
pub use response_processing::GenerateResponseProcessingStage;
|
||||
pub(crate) use preparation::GeneratePreparationStage;
|
||||
pub(crate) use request_building::GenerateRequestBuildingStage;
|
||||
pub(crate) use response_processing::GenerateResponseProcessingStage;
|
||||
|
||||
@@ -23,7 +23,7 @@ use crate::{
|
||||
///
|
||||
/// Extracts generate-specific preparation logic from the old unified PreparationStage.
|
||||
/// This is a direct extraction without architectural changes.
|
||||
pub struct GeneratePreparationStage;
|
||||
pub(crate) struct GeneratePreparationStage;
|
||||
|
||||
#[async_trait]
|
||||
impl PipelineStage for GeneratePreparationStage {
|
||||
|
||||
@@ -18,7 +18,7 @@ use crate::routers::{
|
||||
/// Generate request building stage
|
||||
///
|
||||
/// Extracts generate-specific request building logic from the old unified RequestBuildingStage.
|
||||
pub struct GenerateRequestBuildingStage {
|
||||
pub(crate) struct GenerateRequestBuildingStage {
|
||||
inject_pd_metadata: bool,
|
||||
}
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ use crate::routers::{
|
||||
/// Generate response processing stage
|
||||
///
|
||||
/// Extracts generate-specific response processing logic from the old unified ResponseProcessingStage.
|
||||
pub struct GenerateResponseProcessingStage {
|
||||
pub(crate) struct GenerateResponseProcessingStage {
|
||||
processor: processor::ResponseProcessor,
|
||||
streaming_processor: Arc<streaming::StreamingProcessor>,
|
||||
}
|
||||
|
||||
@@ -2,19 +2,15 @@
|
||||
//!
|
||||
//! This module defines stages specific to regular tokenizer-based models.
|
||||
|
||||
pub mod chat;
|
||||
pub mod classify;
|
||||
pub mod embedding;
|
||||
pub mod generate;
|
||||
mod preparation;
|
||||
mod request_building;
|
||||
mod response_processing;
|
||||
pub(crate) mod chat;
|
||||
pub(crate) mod classify;
|
||||
pub(crate) mod embedding;
|
||||
pub(crate) mod generate;
|
||||
pub(crate) mod preparation;
|
||||
pub(crate) mod request_building;
|
||||
pub(crate) mod response_processing;
|
||||
|
||||
pub use chat::{ChatPreparationStage, ChatRequestBuildingStage, ChatResponseProcessingStage};
|
||||
pub use classify::ClassifyResponseProcessingStage;
|
||||
pub use generate::{
|
||||
GeneratePreparationStage, GenerateRequestBuildingStage, GenerateResponseProcessingStage,
|
||||
};
|
||||
pub use preparation::PreparationStage;
|
||||
pub use request_building::RequestBuildingStage;
|
||||
pub use response_processing::ResponseProcessingStage;
|
||||
// Re-export main stages used by pipeline
|
||||
pub(crate) use preparation::PreparationStage;
|
||||
pub(crate) use request_building::RequestBuildingStage;
|
||||
pub(crate) use response_processing::ResponseProcessingStage;
|
||||
|
||||
@@ -20,7 +20,7 @@ use crate::routers::{
|
||||
};
|
||||
|
||||
/// Preparation stage (delegates to endpoint-specific implementations)
|
||||
pub struct PreparationStage {
|
||||
pub(crate) struct PreparationStage {
|
||||
chat_stage: ChatPreparationStage,
|
||||
generate_stage: GeneratePreparationStage,
|
||||
embedding_stage: EmbeddingPreparationStage,
|
||||
|
||||
@@ -17,7 +17,7 @@ use crate::routers::{
|
||||
};
|
||||
|
||||
/// Request building stage (delegates to endpoint-specific implementations)
|
||||
pub struct RequestBuildingStage {
|
||||
pub(crate) struct RequestBuildingStage {
|
||||
chat_stage: ChatRequestBuildingStage,
|
||||
generate_stage: GenerateRequestBuildingStage,
|
||||
embedding_stage: EmbeddingRequestBuildingStage,
|
||||
|
||||
@@ -21,7 +21,7 @@ use crate::routers::{
|
||||
};
|
||||
|
||||
/// Response processing stage (delegates to endpoint-specific implementations)
|
||||
pub struct ResponseProcessingStage {
|
||||
pub(crate) struct ResponseProcessingStage {
|
||||
chat_stage: ChatResponseProcessingStage,
|
||||
generate_stage: GenerateResponseProcessingStage,
|
||||
embedding_stage: EmbeddingResponseProcessingStage,
|
||||
|
||||
@@ -38,7 +38,7 @@ use crate::{
|
||||
|
||||
/// Shared streaming processor for both single and dual dispatch modes
|
||||
#[derive(Clone)]
|
||||
pub struct StreamingProcessor {
|
||||
pub(crate) struct StreamingProcessor {
|
||||
tool_parser_factory: ToolParserFactory,
|
||||
reasoning_parser_factory: ReasoningParserFactory,
|
||||
configured_tool_parser: Option<String>,
|
||||
@@ -1324,7 +1324,9 @@ impl StreamingProcessor {
|
||||
}
|
||||
|
||||
/// Build SSE response with proper headers
|
||||
pub fn build_sse_response(rx: mpsc::UnboundedReceiver<Result<Bytes, io::Error>>) -> Response {
|
||||
pub(crate) fn build_sse_response(
|
||||
rx: mpsc::UnboundedReceiver<Result<Bytes, io::Error>>,
|
||||
) -> Response {
|
||||
let stream = UnboundedReceiverStream::new(rx);
|
||||
let mut response = Response::new(Body::from_stream(stream));
|
||||
*response.status_mut() = StatusCode::OK;
|
||||
|
||||
@@ -49,7 +49,7 @@ use crate::{
|
||||
/// preparation stages (chat, generate, embedding).
|
||||
///
|
||||
/// Returns the tokenizer Arc, which is also cached in `ctx.state.tokenizer`.
|
||||
pub fn resolve_tokenizer(
|
||||
pub(crate) fn resolve_tokenizer(
|
||||
ctx: &mut RequestContext,
|
||||
stage_name: &str,
|
||||
) -> Result<Arc<dyn Tokenizer>, Box<Response>> {
|
||||
@@ -87,7 +87,9 @@ pub fn resolve_tokenizer(
|
||||
}
|
||||
|
||||
/// Get gRPC client from worker, returning appropriate error response on failure
|
||||
pub async fn get_grpc_client_from_worker(worker: &Arc<dyn Worker>) -> Result<GrpcClient, Response> {
|
||||
pub(crate) async fn get_grpc_client_from_worker(
|
||||
worker: &Arc<dyn Worker>,
|
||||
) -> Result<GrpcClient, Response> {
|
||||
// Get cached client from worker (or create one if not cached yet)
|
||||
let client_arc = worker
|
||||
.get_grpc_client()
|
||||
@@ -157,7 +159,7 @@ fn process_tool_call_arguments(messages: &mut [Value]) -> Result<(), String> {
|
||||
}
|
||||
|
||||
/// Process messages based on content format for ANY message type
|
||||
pub fn process_content_format(
|
||||
pub(crate) fn process_content_format(
|
||||
messages: &[ChatMessage],
|
||||
content_format: ChatTemplateContentFormat,
|
||||
) -> Result<Vec<Value>, String> {
|
||||
@@ -227,7 +229,7 @@ fn transform_content_field(content_value: &mut Value, content_format: ChatTempla
|
||||
|
||||
/// Generate tool constraints for structured generation
|
||||
/// Note: tools should already be filtered if needed (by allowed_tools or specific function)
|
||||
pub fn generate_tool_constraints(
|
||||
pub(crate) fn generate_tool_constraints(
|
||||
tools: &[Tool],
|
||||
tool_choice: &Option<ToolChoice>,
|
||||
_model: &str,
|
||||
@@ -343,7 +345,7 @@ fn build_required_array_schema(tools: &[Tool]) -> Result<String, String> {
|
||||
///
|
||||
/// Returns filtered tools if filtering is needed, otherwise returns None.
|
||||
/// Used by both Chat API and Responses API (Harmony) for constraint generation.
|
||||
pub fn filter_tools_by_tool_choice(
|
||||
pub(crate) fn filter_tools_by_tool_choice(
|
||||
tools: &[Tool],
|
||||
tool_choice: &Option<ToolChoice>,
|
||||
) -> Option<Vec<Tool>> {
|
||||
@@ -377,7 +379,7 @@ pub fn filter_tools_by_tool_choice(
|
||||
///
|
||||
/// Note: Tool existence is validated earlier in ChatCompletionRequest::validate(),
|
||||
/// so this function assumes tool_choice references valid tools.
|
||||
pub fn filter_chat_request_by_tool_choice(
|
||||
pub(crate) fn filter_chat_request_by_tool_choice(
|
||||
body: &ChatCompletionRequest,
|
||||
) -> std::borrow::Cow<'_, ChatCompletionRequest> {
|
||||
if let Some(tools) = &body.tools {
|
||||
@@ -394,7 +396,7 @@ pub fn filter_chat_request_by_tool_choice(
|
||||
|
||||
/// Process chat messages and apply template (shared by both routers)
|
||||
/// Requires HuggingFace tokenizer with chat template support
|
||||
pub fn process_chat_messages(
|
||||
pub(crate) fn process_chat_messages(
|
||||
request: &ChatCompletionRequest,
|
||||
tokenizer: &dyn Tokenizer,
|
||||
) -> Result<ProcessedMessages, String> {
|
||||
@@ -515,7 +517,7 @@ pub fn process_chat_messages(
|
||||
}
|
||||
|
||||
/// Create a StopSequenceDecoder from stop parameters
|
||||
pub fn create_stop_decoder(
|
||||
pub(crate) fn create_stop_decoder(
|
||||
tokenizer: &Arc<dyn Tokenizer>,
|
||||
stop: Option<&StringOrArray>,
|
||||
stop_token_ids: Option<&Vec<u32>>,
|
||||
@@ -557,7 +559,7 @@ pub fn create_stop_decoder(
|
||||
}
|
||||
|
||||
/// Parse tool calls from JSON schema constrained response
|
||||
pub fn parse_json_schema_response(
|
||||
pub(crate) fn parse_json_schema_response(
|
||||
processed_text: &str,
|
||||
tool_choice: &Option<ToolChoice>,
|
||||
model: &str,
|
||||
@@ -646,7 +648,7 @@ pub fn parse_json_schema_response(
|
||||
/// # Returns
|
||||
/// * `Ok(Vec<GenerateComplete>)` - All complete responses collected from the stream
|
||||
/// * `Err(Response)` - Error response if the stream fails or returns an error
|
||||
pub async fn collect_stream_responses(
|
||||
pub(crate) async fn collect_stream_responses(
|
||||
stream: &mut ProtoStream,
|
||||
worker_name: &str,
|
||||
) -> Result<Vec<ProtoGenerateComplete>, Response> {
|
||||
@@ -691,7 +693,7 @@ pub async fn collect_stream_responses(
|
||||
|
||||
/// Count the number of tool calls in the request message history
|
||||
/// This is used for KimiK2 format which needs globally unique indices
|
||||
pub fn get_history_tool_calls_count(request: &ChatCompletionRequest) -> usize {
|
||||
pub(crate) fn get_history_tool_calls_count(request: &ChatCompletionRequest) -> usize {
|
||||
request
|
||||
.messages
|
||||
.iter()
|
||||
@@ -715,7 +717,7 @@ pub fn get_history_tool_calls_count(request: &ChatCompletionRequest) -> usize {
|
||||
///
|
||||
/// # Returns
|
||||
/// A unique ID string. KimiK2 uses `functions.{name}:{global_index}`, others use `call_{uuid}`
|
||||
pub fn generate_tool_call_id(
|
||||
pub(crate) fn generate_tool_call_id(
|
||||
model: &str,
|
||||
tool_name: &str,
|
||||
tool_index: usize,
|
||||
@@ -737,7 +739,7 @@ pub fn generate_tool_call_id(
|
||||
}
|
||||
|
||||
/// Check if a reasoning parser is available for the given model
|
||||
pub fn check_reasoning_parser_availability(
|
||||
pub(crate) fn check_reasoning_parser_availability(
|
||||
reasoning_parser_factory: &ReasoningParserFactory,
|
||||
configured_parser: Option<&str>,
|
||||
model: &str,
|
||||
@@ -752,7 +754,7 @@ pub fn check_reasoning_parser_availability(
|
||||
}
|
||||
|
||||
/// Check if a tool parser is available for the given model
|
||||
pub fn check_tool_parser_availability(
|
||||
pub(crate) fn check_tool_parser_availability(
|
||||
tool_parser_factory: &ToolParserFactory,
|
||||
configured_parser: Option<&str>,
|
||||
model: &str,
|
||||
@@ -769,7 +771,7 @@ pub fn check_tool_parser_availability(
|
||||
/// If a parser name is explicitly configured, use that parser.
|
||||
/// Otherwise, auto-detect based on the model name.
|
||||
/// Get a pooled reasoning parser (for non-streaming where state doesn't matter)
|
||||
pub fn get_reasoning_parser(
|
||||
pub(crate) fn get_reasoning_parser(
|
||||
reasoning_parser_factory: &ReasoningParserFactory,
|
||||
configured_parser: Option<&str>,
|
||||
model: &str,
|
||||
@@ -793,7 +795,7 @@ pub fn get_reasoning_parser(
|
||||
}
|
||||
|
||||
/// Create a fresh reasoning parser instance (for streaming where state isolation is needed)
|
||||
pub fn create_reasoning_parser(
|
||||
pub(crate) fn create_reasoning_parser(
|
||||
reasoning_parser_factory: &ReasoningParserFactory,
|
||||
configured_parser: Option<&str>,
|
||||
model: &str,
|
||||
@@ -821,7 +823,7 @@ pub fn create_reasoning_parser(
|
||||
/// If a parser name is explicitly configured, use that parser.
|
||||
/// Otherwise, auto-detect based on the model name.
|
||||
/// Get a pooled tool parser (for non-streaming where state doesn't matter)
|
||||
pub fn get_tool_parser(
|
||||
pub(crate) fn get_tool_parser(
|
||||
tool_parser_factory: &ToolParserFactory,
|
||||
configured_parser: Option<&str>,
|
||||
model: &str,
|
||||
@@ -845,7 +847,7 @@ pub fn get_tool_parser(
|
||||
}
|
||||
|
||||
/// Create a fresh tool parser instance (for streaming where state isolation is needed)
|
||||
pub fn create_tool_parser(
|
||||
pub(crate) fn create_tool_parser(
|
||||
tool_parser_factory: &ToolParserFactory,
|
||||
configured_parser: Option<&str>,
|
||||
model: &str,
|
||||
@@ -872,7 +874,7 @@ pub fn create_tool_parser(
|
||||
///
|
||||
/// This function decodes token IDs using the tokenizer and builds the logprobs structure
|
||||
/// expected by the OpenAI API format.
|
||||
pub fn convert_proto_to_openai_logprobs(
|
||||
pub(crate) fn convert_proto_to_openai_logprobs(
|
||||
proto_logprobs: &OutputLogProbs,
|
||||
tokenizer: &Arc<dyn Tokenizer>,
|
||||
) -> Result<ChatLogProbs, String> {
|
||||
@@ -949,7 +951,9 @@ pub fn convert_proto_to_openai_logprobs(
|
||||
///
|
||||
/// Generate format: [[logprob, token_id, ...], [logprob, token_id, ...], ...]
|
||||
/// Each inner vec contains [logprob (f64), token_id (i32), ...]
|
||||
pub fn convert_generate_output_logprobs(proto_logprobs: &OutputLogProbs) -> Vec<Vec<Option<f64>>> {
|
||||
pub(crate) fn convert_generate_output_logprobs(
|
||||
proto_logprobs: &OutputLogProbs,
|
||||
) -> Vec<Vec<Option<f64>>> {
|
||||
proto_logprobs
|
||||
.token_logprobs
|
||||
.iter()
|
||||
@@ -962,7 +966,9 @@ pub fn convert_generate_output_logprobs(proto_logprobs: &OutputLogProbs) -> Vec<
|
||||
///
|
||||
/// Generate format: [[logprob, token_id, ...], [logprob, token_id, ...], ...]
|
||||
/// First token has null logprob: [[null, token_id], [logprob, token_id], ...]
|
||||
pub fn convert_generate_input_logprobs(proto_logprobs: &InputLogProbs) -> Vec<Vec<Option<f64>>> {
|
||||
pub(crate) fn convert_generate_input_logprobs(
|
||||
proto_logprobs: &InputLogProbs,
|
||||
) -> Vec<Vec<Option<f64>>> {
|
||||
proto_logprobs
|
||||
.token_logprobs
|
||||
.iter()
|
||||
@@ -985,7 +991,10 @@ pub fn convert_generate_input_logprobs(proto_logprobs: &InputLogProbs) -> Vec<Ve
|
||||
/// - Any other JSON -> Other(...)
|
||||
///
|
||||
/// For backward compatibility, also handles simple string "stop" -> Stop
|
||||
pub fn parse_finish_reason(reason_str: &str, completion_tokens: i32) -> GenerateFinishReason {
|
||||
pub(crate) fn parse_finish_reason(
|
||||
reason_str: &str,
|
||||
completion_tokens: i32,
|
||||
) -> GenerateFinishReason {
|
||||
if reason_str == "stop" {
|
||||
return GenerateFinishReason::Stop;
|
||||
}
|
||||
@@ -1010,7 +1019,7 @@ pub fn parse_finish_reason(reason_str: &str, completion_tokens: i32) -> Generate
|
||||
// ============================================================================
|
||||
|
||||
/// Map route path to endpoint label for metrics
|
||||
pub fn route_to_endpoint(route: &str) -> &'static str {
|
||||
pub(crate) fn route_to_endpoint(route: &str) -> &'static str {
|
||||
match route {
|
||||
"/v1/chat/completions" => metrics_labels::ENDPOINT_CHAT,
|
||||
"/generate" => metrics_labels::ENDPOINT_GENERATE,
|
||||
@@ -1022,7 +1031,7 @@ pub fn route_to_endpoint(route: &str) -> &'static str {
|
||||
}
|
||||
|
||||
/// Map HTTP status code to error type label for metrics
|
||||
pub fn error_type_from_status(status: StatusCode) -> &'static str {
|
||||
pub(crate) fn error_type_from_status(status: StatusCode) -> &'static str {
|
||||
match status.as_u16() {
|
||||
400 => metrics_labels::ERROR_VALIDATION,
|
||||
404 => metrics_labels::ERROR_NO_WORKERS,
|
||||
|
||||
Reference in New Issue
Block a user