use std::fs::{self, File};
use std::io::{self, Read};
use std::path::{Path, PathBuf};
use base64::engine::general_purpose::STANDARD;
use base64::Engine;
use crate::canonical::{
CanonicalError, CanonicalRequest, Content, DocumentSource, ErrorKind, ImageSource, Message,
Role,
};
use crate::pipeline::parse::parse;
pub fn open_input(path: Option<&Path>) -> io::Result<Box<dyn Read>> {
Ok(match path {
Some(path) => Box::new(File::open(path)?),
None => Box::new(io::stdin().lock()),
})
}
const STDIN_PATH: &str = "-";
fn media_type(path: &Path) -> Option<&'static str> {
match path.extension()?.to_str()?.to_ascii_lowercase().as_str() {
"png" => Some("image/png"),
"jpg" | "jpeg" => Some("image/jpeg"),
"gif" => Some("image/gif"),
"webp" => Some("image/webp"),
"pdf" => Some("application/pdf"),
_ => None,
}
}
fn media_part(media_type: &str, bytes: Vec<u8>) -> Content {
let (media_type, data) = (media_type.to_owned(), STANDARD.encode(bytes));
match media_type.starts_with("image/") {
true => Content::Image {
source: ImageSource::Base64 { media_type, data },
},
false => Content::Document {
source: DocumentSource::Base64 { media_type, data },
},
}
}
pub fn read_files(
paths: &[PathBuf],
stdin: &mut dyn Read,
) -> Result<Vec<Content>, (PathBuf, io::Error)> {
paths
.iter()
.map(|p| {
let mut buf = String::new();
match (p.as_os_str() == STDIN_PATH, media_type(p)) {
(true, _) => stdin.read_to_string(&mut buf).map(|_| Content::Text(buf)),
(false, Some(mt)) => fs::read(p).map(|bytes| media_part(mt, bytes)),
(false, None) => fs::read_to_string(p).map(Content::Text),
}
.map_err(|e| (p.clone(), e))
})
.collect()
}
pub fn read_request(
prompt: Option<&str>,
files: Vec<Content>,
reader: &mut dyn Read,
) -> Result<CanonicalRequest, CanonicalError> {
match prompt {
Some(prompt) => Ok(user_message(push_text(files, prompt))),
None if files.is_empty() => parse(reader),
None => {
let mut buf = Vec::new();
reader.read_to_end(&mut buf).map_err(read_err)?;
if buf.iter().any(|b| !b.is_ascii_whitespace()) {
Err(cannot_combine())
} else {
Ok(user_message(files))
}
}
}
}
fn push_text(mut parts: Vec<Content>, prompt: &str) -> Vec<Content> {
parts.push(Content::Text(prompt.to_owned()));
parts
}
fn user_message(content: Vec<Content>) -> CanonicalRequest {
CanonicalRequest {
messages: vec![Message {
role: Role::User,
content,
}],
..Default::default()
}
}
fn cannot_combine() -> CanonicalError {
CanonicalError {
kind: ErrorKind::Usage,
message: "cannot combine --file with a canonical request on stdin \
(put the file contents in the request's messages instead)"
.to_owned(),
provider_detail: None,
retry_after_seconds: None,
}
}
fn read_err(e: io::Error) -> CanonicalError {
CanonicalError {
kind: ErrorKind::ParseInput,
message: format!("failed to read stdin: {e}"),
provider_detail: None,
retry_after_seconds: None,
}
}