Skip to main content

elph_ai/api/
faux.rs

1use std::collections::HashMap;
2use std::sync::{Arc, Mutex};
3
4use serde_json::Value;
5
6use crate::types::ToolCall;
7use crate::types::{AssistantContentBlock, AssistantMessage, AssistantMessageEvent, Context, Model, ProviderResponse};
8use crate::types::{ProviderStreams, SimpleStreamOptions, StopReason, StreamOptions, TextContent, ThinkingContent};
9use crate::utils::event_stream::AssistantMessageEventStream;
10
11const DEFAULT_API: &str = "faux";
12const DEFAULT_PROVIDER: &str = "faux";
13const DEFAULT_MODEL_ID: &str = "faux-1";
14const DEFAULT_BASE_URL: &str = "http://localhost:0";
15const DEFAULT_MIN_TOKEN_SIZE: usize = 3;
16const DEFAULT_MAX_TOKEN_SIZE: usize = 5;
17const MAX_PROMPT_CACHE_ENTRIES: usize = 128;
18
19pub type FauxResponseFactory =
20    Arc<dyn Fn(&Context, Option<&StreamOptions>, &FauxState, &Model) -> AssistantMessage + Send + Sync>;
21
22#[allow(clippy::large_enum_variant)]
23pub enum FauxResponseStep {
24    Static(AssistantMessage),
25    Factory(FauxResponseFactory),
26}
27
28#[derive(Debug, Default, Clone)]
29pub struct FauxState {
30    pub call_count: u64,
31}
32
33#[derive(Default)]
34pub struct RegisterFauxProviderOptions {
35    pub api: Option<String>,
36    pub provider: Option<String>,
37    pub models: Option<Vec<FauxModelDefinition>>,
38    pub tokens_per_second: Option<f64>,
39    pub token_size_min: Option<usize>,
40    pub token_size_max: Option<usize>,
41}
42
43#[derive(Clone)]
44pub struct FauxModelDefinition {
45    pub id: String,
46    pub name: Option<String>,
47    pub reasoning: Option<bool>,
48    pub input: Option<Vec<String>>,
49    pub context_window: Option<u32>,
50    pub max_tokens: Option<u32>,
51}
52
53pub struct FauxCore {
54    pub api: String,
55    pub provider: String,
56    pub models: Vec<Model>,
57    pub state: Arc<Mutex<FauxState>>,
58    pending: Arc<Mutex<Vec<FauxResponseStep>>>,
59    tokens_per_second: Option<f64>,
60    min_token_size: usize,
61    max_token_size: usize,
62    prompt_cache: Arc<Mutex<HashMap<String, String>>>,
63}
64
65pub struct FauxApi {
66    core: Arc<FauxCore>,
67}
68
69impl FauxCore {
70    pub fn new(options: RegisterFauxProviderOptions) -> Self {
71        let min = options.token_size_min.unwrap_or(DEFAULT_MIN_TOKEN_SIZE).max(1);
72        let max = options.token_size_max.unwrap_or(DEFAULT_MAX_TOKEN_SIZE).max(min);
73        let api = options.api.unwrap_or_else(|| DEFAULT_API.to_string());
74        let provider = options.provider.unwrap_or_else(|| DEFAULT_PROVIDER.to_string());
75        let defs = options.models.unwrap_or_else(|| {
76            vec![FauxModelDefinition {
77                id: DEFAULT_MODEL_ID.to_string(),
78                name: Some("Faux Model".to_string()),
79                reasoning: Some(false),
80                input: Some(vec!["text".to_string(), "image".to_string()]),
81                context_window: Some(128_000),
82                max_tokens: Some(16_384),
83            }]
84        });
85        let models = defs
86            .into_iter()
87            .map(|d| Model {
88                id: d.id.clone(),
89                name: d.name.unwrap_or_else(|| d.id.clone()),
90                api: api.clone(),
91                provider: provider.clone(),
92                base_url: DEFAULT_BASE_URL.to_string(),
93                reasoning: d.reasoning.unwrap_or(false),
94                thinking_level_map: None,
95                input: d.input.unwrap_or_else(|| vec!["text".to_string(), "image".to_string()]),
96                cost: crate::types::ModelCost {
97                    input: 0.0,
98                    output: 0.0,
99                    cache_read: 0.0,
100                    cache_write: 0.0,
101
102                    tiers: None,
103                },
104                context_window: d.context_window.unwrap_or(128_000),
105                max_tokens: d.max_tokens.unwrap_or(16_384),
106                headers: None,
107                openai_completions_compat: None,
108                openai_responses_compat: None,
109                anthropic_compat: None,
110            })
111            .collect();
112
113        Self {
114            api,
115            provider,
116            models,
117            state: Arc::new(Mutex::new(FauxState::default())),
118            pending: Arc::new(Mutex::new(Vec::new())),
119            tokens_per_second: options.tokens_per_second,
120            min_token_size: min,
121            max_token_size: max,
122            prompt_cache: Arc::new(Mutex::new(HashMap::new())),
123        }
124    }
125
126    pub fn set_responses(&self, responses: Vec<FauxResponseStep>) {
127        *self.pending.lock().unwrap() = responses;
128    }
129
130    pub fn append_responses(&self, responses: Vec<FauxResponseStep>) {
131        self.pending.lock().unwrap().extend(responses);
132    }
133
134    pub fn pending_count(&self) -> usize {
135        self.pending.lock().unwrap().len()
136    }
137
138    pub fn api(self: &Arc<Self>) -> FauxApi {
139        FauxApi { core: Arc::clone(self) }
140    }
141}
142
143impl ProviderStreams for FauxApi {
144    fn stream(&self, model: &Model, context: &Context, options: Option<StreamOptions>) -> AssistantMessageEventStream {
145        self.stream_simple(model, context, options.map(SimpleStreamOptions::from_stream))
146    }
147
148    fn stream_simple(
149        &self,
150        model: &Model,
151        context: &Context,
152        options: Option<SimpleStreamOptions>,
153    ) -> AssistantMessageEventStream {
154        let stream = AssistantMessageEventStream::new();
155        let core = self.core.clone();
156        let model = model.clone();
157        let context = context.clone();
158        let options = options.map(|o| o.base);
159        let s = stream.clone();
160        tokio::spawn(async move {
161            if let Err(e) = run_faux(&core, &model, &context, options.as_ref(), &s).await {
162                let mut output = AssistantMessage::empty(&model);
163                output.stop_reason = StopReason::Error;
164                output.error_message = Some(e.to_string());
165                s.push(AssistantMessageEvent::Error {
166                    reason: StopReason::Error,
167                    error: output.clone(),
168                });
169                s.end();
170            }
171        });
172        stream
173    }
174}
175
176async fn run_faux(
177    core: &FauxCore,
178    model: &Model,
179    context: &Context,
180    options: Option<&StreamOptions>,
181    stream: &AssistantMessageEventStream,
182) -> anyhow::Result<()> {
183    if crate::api::common::is_request_aborted(&options.and_then(|o| o.signal.clone())) {
184        let mut output = AssistantMessage::empty(model);
185        crate::api::common::finish_stream_error(stream, &mut output, crate::api::common::request_aborted_error(), true);
186        return Ok(());
187    }
188    {
189        let mut state = core.state.lock().unwrap();
190        state.call_count += 1;
191    }
192
193    let step = {
194        let mut pending = core.pending.lock().unwrap();
195        if pending.is_empty() {
196            None
197        } else {
198            Some(pending.remove(0))
199        }
200    };
201    let state = core.state.lock().unwrap().clone();
202
203    let message = match step {
204        Some(FauxResponseStep::Static(m)) => m,
205        Some(FauxResponseStep::Factory(f)) => f(context, options, &state, model),
206        None => {
207            let mut m = AssistantMessage::empty(model);
208            m.stop_reason = StopReason::Error;
209            m.error_message = Some("No more faux responses queued".to_string());
210            m
211        }
212    };
213
214    crate::api::common::apply_on_response(
215        options.and_then(|o| o.on_response.as_ref()),
216        ProviderResponse {
217            status: 200,
218            headers: HashMap::from([
219                ("x-faux-provider".to_string(), "ok".to_string()),
220                ("content-type".to_string(), "text/event-stream".to_string()),
221            ]),
222        },
223        model,
224    )
225    .await;
226
227    let message = with_usage_estimate(message, context, options, &core.prompt_cache);
228    stream_with_deltas(
229        stream,
230        message,
231        core.min_token_size,
232        core.max_token_size,
233        core.tokens_per_second,
234        options.and_then(|o| o.signal.clone()),
235    )
236    .await;
237    Ok(())
238}
239
240pub fn faux_text(text: impl Into<String>) -> AssistantContentBlock {
241    AssistantContentBlock::Text(TextContent::new(text))
242}
243
244pub fn faux_thinking(thinking: impl Into<String>) -> AssistantContentBlock {
245    AssistantContentBlock::Thinking(ThinkingContent::new(thinking))
246}
247
248pub fn faux_tool_call(name: impl Into<String>, arguments: Value, id: Option<String>) -> AssistantContentBlock {
249    AssistantContentBlock::ToolCall(ToolCall::new(
250        id.unwrap_or_else(|| format!("tool:{}", chrono::Utc::now().timestamp_millis())),
251        name,
252        arguments,
253    ))
254}
255
256pub fn faux_assistant_message(
257    content: Vec<AssistantContentBlock>,
258    stop_reason: Option<StopReason>,
259) -> AssistantMessage {
260    let mut m = AssistantMessage::empty(&Model {
261        id: DEFAULT_MODEL_ID.to_string(),
262        name: DEFAULT_MODEL_ID.to_string(),
263        api: DEFAULT_API.to_string(),
264        provider: DEFAULT_PROVIDER.to_string(),
265        base_url: DEFAULT_BASE_URL.to_string(),
266        reasoning: false,
267        thinking_level_map: None,
268        input: vec!["text".to_string()],
269        cost: crate::types::ModelCost {
270            input: 0.0,
271            output: 0.0,
272            cache_read: 0.0,
273            cache_write: 0.0,
274
275            tiers: None,
276        },
277        context_window: 128_000,
278        max_tokens: 16_384,
279        headers: None,
280        openai_completions_compat: None,
281        openai_responses_compat: None,
282        anthropic_compat: None,
283    });
284    m.content = content;
285    m.stop_reason = stop_reason.unwrap_or(StopReason::Stop);
286    m
287}
288
289fn estimate_tokens(text: &str) -> u64 {
290    ((text.len() as f64) / 4.0).ceil() as u64
291}
292
293fn with_usage_estimate(
294    mut message: AssistantMessage,
295    context: &Context,
296    options: Option<&StreamOptions>,
297    prompt_cache: &Arc<Mutex<HashMap<String, String>>>,
298) -> AssistantMessage {
299    let prompt_text = serialize_context(context);
300    let prompt_tokens = estimate_tokens(&prompt_text);
301    let output_text: String = message
302        .content
303        .iter()
304        .map(|b| match b {
305            AssistantContentBlock::Text(t) => t.text.clone(),
306            AssistantContentBlock::Thinking(t) => t.thinking.clone(),
307            AssistantContentBlock::ToolCall(tc) => format!("{}:{}", tc.name, tc.arguments),
308        })
309        .collect::<Vec<_>>()
310        .join("");
311    let output_tokens = estimate_tokens(&output_text);
312    let mut input = prompt_tokens;
313    let mut cache_read = 0u64;
314    let mut cache_write = 0u64;
315
316    if let Some(session_id) = options.and_then(|o| o.session_id.as_deref()) {
317        let mut cache = prompt_cache.lock().unwrap();
318        if cache.len() >= MAX_PROMPT_CACHE_ENTRIES
319            && !cache.contains_key(session_id)
320            && let Some(evicted) = cache.keys().next().cloned()
321        {
322            cache.remove(&evicted);
323        }
324        if let Some(previous) = cache.get(session_id).cloned() {
325            let cached_chars = common_prefix_len(&previous, &prompt_text);
326            cache_read = estimate_tokens(&previous[..cached_chars.min(previous.len())]);
327            cache_write = estimate_tokens(&prompt_text[cached_chars.min(prompt_text.len())..]);
328            input = prompt_tokens.saturating_sub(cache_read);
329        } else {
330            cache_write = prompt_tokens;
331        }
332        cache.insert(session_id.to_string(), prompt_text);
333    }
334
335    message.usage.input = input;
336    message.usage.output = output_tokens;
337    message.usage.cache_read = cache_read;
338    message.usage.cache_write = cache_write;
339    message.usage.total_tokens = input + output_tokens + cache_read + cache_write;
340    message
341}
342
343fn serialize_context(context: &Context) -> String {
344    let mut parts = Vec::new();
345    if let Some(sp) = &context.system_prompt {
346        parts.push(format!("system:{sp}"));
347    }
348    for msg in &context.messages {
349        parts.push(format!("{msg:?}"));
350    }
351    if let Some(tools) = &context.tools {
352        parts.push(format!("tools:{}", serde_json::to_string(tools).unwrap_or_default()));
353    }
354    parts.join("\n\n")
355}
356
357fn common_prefix_len(a: &str, b: &str) -> usize {
358    a.chars().zip(b.chars()).take_while(|(x, y)| x == y).count()
359}
360
361fn split_by_token_size(text: &str, min: usize, max: usize) -> Vec<String> {
362    let mut chunks = Vec::new();
363    let mut index = 0;
364    let bytes = text.as_bytes();
365    while index < bytes.len() {
366        let token_size = {
367            use rand::RngExt;
368            rand::rng().random_range(min..=max)
369        };
370        let char_size = (token_size * 4).max(1);
371        let end = (index + char_size).min(bytes.len());
372        chunks.push(String::from_utf8_lossy(&bytes[index..end]).to_string());
373        index = end;
374    }
375    if chunks.is_empty() {
376        chunks.push(String::new());
377    }
378    chunks
379}
380
381async fn schedule_chunk(chunk: &str, tokens_per_second: Option<f64>) {
382    if let Some(tps) = tokens_per_second
383        && tps > 0.0
384    {
385        let delay_ms = (estimate_tokens(chunk) as f64 / tps * 1000.0) as u64;
386        tokio::time::sleep(std::time::Duration::from_millis(delay_ms)).await;
387    }
388}
389
390async fn stream_with_deltas(
391    stream: &AssistantMessageEventStream,
392    message: AssistantMessage,
393    min_token_size: usize,
394    max_token_size: usize,
395    tokens_per_second: Option<f64>,
396    signal: Option<tokio_util::sync::CancellationToken>,
397) {
398    let mut partial = AssistantMessage {
399        content: vec![],
400        ..message.clone()
401    };
402    stream.push(AssistantMessageEvent::Start {
403        partial: partial.clone(),
404    });
405
406    for (index, block) in message.content.iter().enumerate() {
407        if crate::api::common::is_request_aborted(&signal) {
408            let mut output = partial.clone();
409            output.stop_reason = StopReason::Aborted;
410            crate::api::common::finish_stream_error(
411                stream,
412                &mut output,
413                crate::api::common::request_aborted_error(),
414                true,
415            );
416            return;
417        }
418        match block {
419            AssistantContentBlock::Thinking(t) => {
420                partial
421                    .content
422                    .push(AssistantContentBlock::Thinking(ThinkingContent::new("")));
423                stream.push(AssistantMessageEvent::ThinkingStart {
424                    content_index: index,
425                    partial: partial.clone(),
426                });
427                for chunk in split_by_token_size(&t.thinking, min_token_size, max_token_size) {
428                    schedule_chunk(&chunk, tokens_per_second).await;
429                    if crate::api::common::is_request_aborted(&signal) {
430                        let mut output = partial.clone();
431                        output.stop_reason = StopReason::Aborted;
432                        crate::api::common::finish_stream_error(
433                            stream,
434                            &mut output,
435                            crate::api::common::request_aborted_error(),
436                            true,
437                        );
438                        return;
439                    }
440                    if let AssistantContentBlock::Thinking(tc) = &mut partial.content[index] {
441                        tc.thinking.push_str(&chunk);
442                    }
443                    stream.push(AssistantMessageEvent::ThinkingDelta {
444                        content_index: index,
445                        delta: chunk,
446                        partial: partial.clone(),
447                    });
448                }
449                stream.push(AssistantMessageEvent::ThinkingEnd {
450                    content_index: index,
451                    content: t.thinking.clone(),
452                    partial: partial.clone(),
453                });
454            }
455            AssistantContentBlock::Text(t) => {
456                partial.content.push(AssistantContentBlock::Text(TextContent::new("")));
457                stream.push(AssistantMessageEvent::TextStart {
458                    content_index: index,
459                    partial: partial.clone(),
460                });
461                for chunk in split_by_token_size(&t.text, min_token_size, max_token_size) {
462                    schedule_chunk(&chunk, tokens_per_second).await;
463                    if crate::api::common::is_request_aborted(&signal) {
464                        let mut output = partial.clone();
465                        output.stop_reason = StopReason::Aborted;
466                        crate::api::common::finish_stream_error(
467                            stream,
468                            &mut output,
469                            crate::api::common::request_aborted_error(),
470                            true,
471                        );
472                        return;
473                    }
474                    if let AssistantContentBlock::Text(tc) = &mut partial.content[index] {
475                        tc.text.push_str(&chunk);
476                    }
477                    stream.push(AssistantMessageEvent::TextDelta {
478                        content_index: index,
479                        delta: chunk,
480                        partial: partial.clone(),
481                    });
482                }
483                stream.push(AssistantMessageEvent::TextEnd {
484                    content_index: index,
485                    content: t.text.clone(),
486                    partial: partial.clone(),
487                });
488            }
489            AssistantContentBlock::ToolCall(tc) => {
490                partial.content.push(AssistantContentBlock::ToolCall(ToolCall::new(
491                    &tc.id,
492                    &tc.name,
493                    Value::Object(Default::default()),
494                )));
495                stream.push(AssistantMessageEvent::ToolcallStart {
496                    content_index: index,
497                    partial: partial.clone(),
498                });
499                let args = tc.arguments.to_string();
500                for chunk in split_by_token_size(&args, min_token_size, max_token_size) {
501                    schedule_chunk(&chunk, tokens_per_second).await;
502                    if crate::api::common::is_request_aborted(&signal) {
503                        let mut output = partial.clone();
504                        output.stop_reason = StopReason::Aborted;
505                        crate::api::common::finish_stream_error(
506                            stream,
507                            &mut output,
508                            crate::api::common::request_aborted_error(),
509                            true,
510                        );
511                        return;
512                    }
513                    stream.push(AssistantMessageEvent::ToolcallDelta {
514                        content_index: index,
515                        delta: chunk,
516                        partial: partial.clone(),
517                    });
518                }
519                if let AssistantContentBlock::ToolCall(slot) = &mut partial.content[index] {
520                    slot.arguments = tc.arguments.clone();
521                }
522                stream.push(AssistantMessageEvent::ToolcallEnd {
523                    content_index: index,
524                    tool_call: tc.clone(),
525                    partial: partial.clone(),
526                });
527            }
528        }
529    }
530
531    if matches!(message.stop_reason, StopReason::Error | StopReason::Aborted) {
532        stream.push(AssistantMessageEvent::Error {
533            reason: message.stop_reason,
534            error: message.clone(),
535        });
536        stream.end();
537        return;
538    }
539
540    stream.push(AssistantMessageEvent::Done {
541        reason: message.stop_reason,
542        message: message.clone(),
543    });
544    stream.end();
545}