use std::str;
use crate::{
args::{ArgValues, FromArgs, StrArg},
bytecode::{CallResult, VM},
defer_drop,
exception_private::{ExcType, RunError, RunResult, SimpleException},
heap::{HeapData, HeapGuard},
os::{MontyPath, OpenCallArgs, OsFunctionCall},
resource::ResourceTracker,
types::{PyTrait, file::FileMode},
value::Value,
};
pub(crate) fn builtin_open(vm: &mut VM<'_, impl ResourceTracker>, args: ArgValues) -> RunResult<CallResult> {
let OpenArgs {
file,
mode,
buffering,
encoding,
errors,
newline,
closefd,
opener,
} = OpenArgs::from_args(args, vm)?;
let mut file = HeapGuard::new(file, vm);
let (file, vm) = file.as_parts_mut();
defer_drop!(mode, vm);
defer_drop!(buffering, vm);
defer_drop!(encoding, vm);
defer_drop!(errors, vm);
defer_drop!(newline, vm);
defer_drop!(closefd, vm);
defer_drop!(opener, vm);
validate_ignored_open_kwarg("buffering", buffering, vm)?;
validate_ignored_open_kwarg("encoding", encoding, vm)?;
validate_ignored_open_kwarg("errors", errors, vm)?;
validate_ignored_open_kwarg("newline", newline, vm)?;
validate_ignored_open_kwarg("closefd", closefd, vm)?;
validate_ignored_open_kwarg("opener", opener, vm)?;
let path = extract_path_string(file, vm)?.to_owned();
let file_mode = mode
.as_ref()
.map_or("r", |m| m.as_str(vm))
.parse::<FileMode>()
.map_err(|e| RunError::from(SimpleException::new_msg(ExcType::ValueError, e)))?;
Ok(CallResult::OsCall(OsFunctionCall::Open(OpenCallArgs {
path: MontyPath::new(path),
mode: file_mode,
})))
}
#[derive(FromArgs)]
#[from_args(name = "open", bad_arg_named)]
struct OpenArgs {
file: Value,
#[from_args(default)]
mode: Option<StrArg>,
#[from_args(default = Value::Int(-1))]
buffering: Value,
#[from_args(default = Value::None)]
encoding: Value,
#[from_args(default = Value::None)]
errors: Value,
#[from_args(default = Value::None)]
newline: Value,
#[from_args(default = Value::Bool(true))]
closefd: Value,
#[from_args(default = Value::None)]
opener: Value,
}
fn extract_path_string<'a>(value: &Value, vm: &'a VM<'_, impl ResourceTracker>) -> RunResult<&'a str> {
let opt = match value {
Value::InternString(string_id) => Some(vm.interns.get_str(*string_id)),
Value::InternBytes(bytes_id) => decode_utf8_path(vm.interns.get_bytes(*bytes_id))?,
Value::Ref(id) => match vm.heap.get(*id) {
HeapData::Str(s) => Some(s.as_str()),
HeapData::Path(p) => Some(p.as_str()),
HeapData::Bytes(b) => decode_utf8_path(b.as_slice())?,
_ => None,
},
_ => None,
};
opt.ok_or_else(|| path_type_error(value, vm))
}
fn decode_utf8_path(bytes: &[u8]) -> RunResult<Option<&str>> {
match str::from_utf8(bytes) {
Ok(s) => Ok(Some(s)),
Err(_) => Err(SimpleException::new_msg(ExcType::UnicodeDecodeError, "can't decode bytes path as UTF-8").into()),
}
}
fn validate_ignored_open_kwarg(name: &str, value: &Value, vm: &VM<'_, impl ResourceTracker>) -> Result<(), RunError> {
let is_default = match name {
"buffering" => matches!(value, Value::Int(-1)),
"encoding" => {
if matches!(value, Value::None) {
true
} else if value.is_str(vm.heap) {
let s = match value {
Value::InternString(id) => vm.interns.get_str(*id),
Value::Ref(id) => match vm.heap.get(*id) {
HeapData::Str(s) => s.as_str(),
_ => "",
},
_ => "",
};
s.eq_ignore_ascii_case("utf-8") || s.eq_ignore_ascii_case("utf8")
} else {
return Err(ExcType::type_error(format!(
"open() argument '{name}' must be str or None, not {}",
value.py_type(vm).cpython_arg_name(vm.heap, vm.interns)
)));
}
}
"errors" | "newline" => {
if matches!(value, Value::None) {
true
} else if value.is_str(vm.heap) {
false
} else {
return Err(ExcType::type_error(format!(
"open() argument '{name}' must be str or None, not {}",
value.py_type(vm).cpython_arg_name(vm.heap, vm.interns)
)));
}
}
"closefd" => matches!(value, Value::Bool(true)),
"opener" => matches!(value, Value::None),
_ => unreachable!("validated open keyword name"),
};
if is_default {
Ok(())
} else {
Err(ExcType::type_error(format!("'{name}' argument is not yet supported")))
}
}
fn path_type_error(value: &Value, vm: &VM<'_, impl ResourceTracker>) -> RunError {
ExcType::type_error(format!(
"expected str, bytes or os.PathLike object, not {}",
value.py_type_name(vm)
))
}