use crate::io::{symlink, ApiResult};
use acorn_core::{prelude::String, util::MimeType, Location, Repository, Scheme};
use acorn_host::{
fs::{file_uri_to_path, SafePath},
http::{HttpClient, HttpPolicy},
source::{self, SourcePolicy},
terminal::Label,
};
use arboard::Clipboard;
use color_eyre::eyre::eyre;
use core::fmt;
use std::{
fs::copy,
path::{Path, PathBuf},
};
use strum::EnumIs;
use tracing::error;
pub enum InputSource {
Location(Location),
Text {
content: String,
origin: TextOrigin,
},
}
#[derive(Clone, Debug, PartialEq, Eq, EnumIs)]
pub enum Source {
Local {
name: Option<String>,
path: PathBuf,
action: Option<SourceAction>,
},
Remote {
name: Option<String>,
identifier: String,
},
Unsupported(String),
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum SourceAction {
#[default]
Reference,
Copy,
Symlink,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum TextOrigin {
Clipboard,
Stdin,
}
impl Source {
pub async fn read(source: &str, offline: bool) -> ApiResult<String> {
match Self::read_context(source, offline) {
| Ok((client, location, policy)) => source::read_text(&client, &location, policy)
.await
.map_err(|why| eyre!("Failed to read source — {why}")),
| Err(why) => Err(why),
}
}
pub async fn read_bytes(source: &str, offline: bool) -> ApiResult<Vec<u8>> {
match Self::read_context(source, offline) {
| Ok((client, location, policy)) => source::read_bytes(&client, &location, policy)
.await
.map_err(|why| eyre!("Failed to read source — {why}")),
| Err(why) => Err(why),
}
}
pub fn parse(source: &str) -> Self {
Self::from_location(Location::from(source), false)
}
fn read_context(source: &str, offline: bool) -> ApiResult<(HttpClient, Location, SourcePolicy)> {
HttpClient::new(HttpPolicy::new())
.map(|client| (client, Location::from(source), SourcePolicy::new().with_offline(offline)))
.map_err(|why| eyre!("Failed to initialize source reader — {why}"))
}
fn from_location(location: Location, simple_as_remote: bool) -> Self {
match location {
| Location::Detailed {
scheme: Scheme::File, uri, ..
} => {
let path = file_uri_to_path(&uri).unwrap_or_else(|_| PathBuf::from(&uri));
Self::Local {
name: None,
path,
action: None,
}
}
| Location::Detailed {
scheme: Scheme::HTTPS | Scheme::HTTP,
uri,
..
} => Self::Remote { name: None, identifier: uri },
| Location::Detailed {
scheme: Scheme::Unsupported,
uri,
..
} if Location::from(uri.as_str()).is_absolute() => Self::Local {
name: None,
path: PathBuf::from(uri),
action: None,
},
| Location::Detailed {
scheme: Scheme::SSH | Scheme::Unsupported,
uri,
..
} if simple_as_remote => Self::Remote { name: None, identifier: uri },
| Location::Detailed {
scheme: Scheme::SSH | Scheme::Unsupported,
uri,
..
} => Self::Unsupported(uri),
| Location::Simple(value) => match Location::from(value.as_str()) {
| parsed @ Location::Detailed { .. } => Self::from_location(parsed, simple_as_remote),
| Location::Simple(_) if simple_as_remote && !Location::from(value.as_str()).is_local() => Self::Remote {
name: Some(value.clone()),
identifier: value,
},
| Location::Simple(_) => Self::Local {
name: None,
path: PathBuf::from(value),
action: None,
},
},
}
}
pub fn name(&self) -> String {
match self {
| Source::Local { name: Some(name), .. } | Source::Remote { name: Some(name), .. } => name.clone(),
| Source::Local { path, .. } => path.file_stem().and_then(|s| s.to_str()).unwrap_or("model").to_string(),
| Source::Remote { identifier, .. } | Source::Unsupported(identifier) => identifier.clone(),
}
}
pub fn identifier(&self) -> String {
match self {
| Source::Local { path, .. } => path.display().to_string(),
| Source::Remote { identifier, .. } | Source::Unsupported(identifier) => identifier.clone(),
}
}
pub fn with_action(self, action: Option<SourceAction>) -> Self {
match self {
| Source::Local { name, path, .. } => Source::Local { name, path, action },
| other => other,
}
}
pub fn with_name(self, value: impl Into<String>) -> Self {
let binding = value.into();
let trimmed = binding.trim();
let name = (!trimmed.is_empty()).then(|| trimmed.to_string());
match self {
| Source::Local { path, action, .. } => Source::Local { name, path, action },
| Source::Remote { identifier, .. } => Source::Remote { name, identifier },
| Source::Unsupported(value) => Source::Unsupported(value),
}
}
}
impl fmt::Display for Source {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let name = self.name();
let identifier = self.identifier();
if name == identifier {
write!(f, "{name}")
} else {
write!(f, "{name} ({identifier})")
}
}
}
impl From<&Repository> for Source {
fn from(repository: &Repository) -> Self {
match repository {
| Repository::HuggingFace { location } => Self::Remote {
name: None,
identifier: repository.id().unwrap_or_else(|| location.uri().unwrap_or_default()),
},
| Repository::OCI { location, .. } => Self::Remote {
name: None,
identifier: location.uri().unwrap_or_default(),
},
| _ => Self::from(repository.location()),
}
}
}
impl From<&str> for Source {
fn from(selector: &str) -> Self {
Self::from_location(Location::from(selector.trim()), true)
}
}
impl From<Location> for Source {
fn from(location: Location) -> Self {
Self::from_location(location, false)
}
}
impl SourceAction {
pub fn from_options(copy: bool, symlink: bool) -> Option<Self> {
match (copy, symlink) {
| (true, false) => Some(Self::Copy),
| (false, true) => Some(Self::Symlink),
| _ => None,
}
}
pub fn materialize(self, path: &Path, root: &Path, target: &SafePath, name: &str) -> ApiResult<String> {
match self {
| SourceAction::Reference => Ok(format!("Local model '{name}' referenced in place at {}", path.display())),
| action => target
.materialize_under(root)
.map_err(Into::<color_eyre::Report>::into)
.and_then(|target| match action {
| SourceAction::Copy if target.exists() => Ok(format!("Skipping copy - '{name}' already exists at {}", target.display())),
| SourceAction::Copy => match copy(path, &target) {
| Ok(_) => Ok(format!("Copied local model '{name}' -> {}", target.display())),
| Err(why) => {
error!("=> {} Copy local model '{}' to {} - {why}", Label::fail(), name, target.display());
Err(why.into())
}
},
| SourceAction::Symlink if target.exists() || target.is_symlink() => {
Ok(format!("Skipping symlink - '{name}' already exists at {}", target.display()))
}
| SourceAction::Symlink => match symlink(path, &target) {
| Ok(_) => Ok(format!("Symlinked local model '{name}' -> {}", target.display())),
| Err(why) => {
error!("=> {} Symlink local model '{}' to {} - {why}", Label::fail(), name, target.display());
Err(why)
}
},
| SourceAction::Reference => Ok(format!("Local model '{name}' referenced in place at {}", path.display())),
}),
}
}
}
impl TextOrigin {
pub fn filename(self, format: &MimeType) -> &'static str {
match (self, format) {
| (Self::Clipboard, MimeType::Cff) => "clipboard.cff",
| (Self::Clipboard, MimeType::Json) => "clipboard.json",
| (Self::Clipboard, MimeType::Jsonc) => "clipboard.jsonc",
| (Self::Clipboard, MimeType::Markdown) => "clipboard.md",
| (Self::Clipboard, MimeType::Yaml) => "clipboard.yaml",
| (Self::Clipboard, MimeType::Zon) => "clipboard.zonf",
| (Self::Clipboard, _) => "clipboard.txt",
| (Self::Stdin, MimeType::Cff) => "stdin.cff",
| (Self::Stdin, MimeType::Json) => "stdin.json",
| (Self::Stdin, MimeType::Jsonc) => "stdin.jsonc",
| (Self::Stdin, MimeType::Markdown) => "stdin.md",
| (Self::Stdin, MimeType::Yaml) => "stdin.yaml",
| (Self::Stdin, MimeType::Zon) => "stdin.zonf",
| (Self::Stdin, _) => "stdin.txt",
}
}
pub fn label(self) -> &'static str {
match self {
| Self::Clipboard => "<clipboard>",
| Self::Stdin => "<stdin>",
}
}
pub fn source(self, content: String) -> ApiResult<InputSource> {
match content.trim().is_empty() {
| true => Err(eyre!("No text received from {}", self.label())),
| false => Ok(InputSource::Text { content, origin: self }),
}
}
}
pub fn read_clipboard() -> ApiResult<String> {
Clipboard::new()
.and_then(|mut clipboard| clipboard.get_text())
.map_err(|why| eyre!("Cannot read clipboard text — {why}"))
}
pub fn select_source<ReadStdin, ReadClipboard>(
path: &Option<PathBuf>,
paste: bool,
filesystem_selected: bool,
watching: bool,
read_stdin: ReadStdin,
read_clipboard: ReadClipboard,
) -> ApiResult<InputSource>
where
ReadStdin: FnOnce() -> Option<String>,
ReadClipboard: FnOnce() -> ApiResult<String>,
{
let location = || {
let input = path
.as_ref()
.map(|path| path.to_string_lossy().to_string())
.unwrap_or_else(|| "./".to_string());
InputSource::Location(Location::from(input.as_str()))
};
match (filesystem_selected, paste) {
| (true, _) => Ok(location()),
| (false, true) => read_clipboard().and_then(|content| TextOrigin::Clipboard.source(content)),
| (false, false) => match read_stdin().filter(|content| !content.trim().is_empty()) {
| Some(_) if watching => Err(eyre!("Standard input cannot be used with --watch")),
| Some(content) => TextOrigin::Stdin.source(content),
| None => Ok(location()),
},
}
}