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