use luna_core::version::LuaVersion;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct Shebang {
pub version: LuaVersion,
}
impl Default for Shebang {
fn default() -> Self {
Shebang {
version: LuaVersion::Lua51,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum ShebangError {
UnknownVersion(String),
}
impl std::fmt::Display for ShebangError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::UnknownVersion(v) => write!(f, "unknown lua version: {v}"),
}
}
}
pub(crate) fn parse(src: &[u8]) -> Result<(Shebang, &[u8]), ShebangError> {
if !src.starts_with(b"#!") {
return Ok((Shebang::default(), src));
}
let line_end = src.iter().position(|&b| b == b'\n').unwrap_or(src.len());
let line = &src[..line_end];
let body = if line_end < src.len() {
&src[line_end + 1..]
} else {
&[]
};
let after_hashbang = &line[2..];
let after_lua = match after_hashbang
.strip_prefix(b"lua")
.or_else(|| after_hashbang.trim_ascii_start().strip_prefix(b"lua"))
{
Some(s) => s,
None => return Ok((Shebang::default(), src)), };
let mut shebang = Shebang::default();
for kv in after_lua
.split(|b| *b == b' ' || *b == b'\t')
.filter(|s| !s.is_empty())
{
if let Some(rest) = kv.strip_prefix(b"version=") {
shebang.version = parse_version(rest)?;
} else if kv.starts_with(b"flags=") || kv.starts_with(b"name=") {
continue;
} else {
continue;
}
}
Ok((shebang, body))
}
fn parse_version(bytes: &[u8]) -> Result<LuaVersion, ShebangError> {
match bytes {
b"5.1" | b"51" => Ok(LuaVersion::Lua51),
b"5.2" | b"52" => Ok(LuaVersion::Lua52),
b"5.3" | b"53" => Ok(LuaVersion::Lua53),
b"5.4" | b"54" => Ok(LuaVersion::Lua54),
b"5.5" | b"55" => Ok(LuaVersion::Lua55),
other => Err(ShebangError::UnknownVersion(
String::from_utf8_lossy(other).into_owned(),
)),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn no_shebang_returns_default_with_src() {
let (s, body) = parse(b"return 1").unwrap();
assert_eq!(s.version, LuaVersion::Lua51);
assert_eq!(body, b"return 1");
}
#[test]
fn shebang_lua_51() {
let (s, body) = parse(b"#!lua version=5.1\nreturn 1").unwrap();
assert_eq!(s.version, LuaVersion::Lua51);
assert_eq!(body, b"return 1");
}
#[test]
fn shebang_lua_53_picks_53() {
let (s, body) = parse(b"#!lua version=5.3\nreturn 1").unwrap();
assert_eq!(s.version, LuaVersion::Lua53);
assert_eq!(body, b"return 1");
}
#[test]
fn shebang_lua_55_picks_55() {
let (s, _) = parse(b"#!lua version=5.5\n").unwrap();
assert_eq!(s.version, LuaVersion::Lua55);
}
#[test]
fn shebang_numeric_form_accepted() {
let (s, _) = parse(b"#!lua version=53\n").unwrap();
assert_eq!(s.version, LuaVersion::Lua53);
}
#[test]
fn shebang_with_extra_keys_tolerated() {
let (s, body) =
parse(b"#!lua version=5.3 flags=no-writes name=mylib\nreturn 1").unwrap();
assert_eq!(s.version, LuaVersion::Lua53);
assert_eq!(body, b"return 1");
}
#[test]
fn shebang_unknown_key_tolerated_forward_compat() {
let (s, body) = parse(b"#!lua version=5.4 future_key=value\nreturn 1").unwrap();
assert_eq!(s.version, LuaVersion::Lua54);
assert_eq!(body, b"return 1");
}
#[test]
fn shebang_unknown_version_rejected() {
let err = parse(b"#!lua version=5.6\nreturn 1").unwrap_err();
assert!(matches!(err, ShebangError::UnknownVersion(ref v) if v == "5.6"));
}
#[test]
fn shebang_without_lua_marker_passes_through() {
let (s, body) = parse(b"#!/foo\nreturn 1").unwrap();
assert_eq!(s.version, LuaVersion::Lua51);
assert_eq!(body, b"#!/foo\nreturn 1");
}
#[test]
fn shebang_with_eof_no_newline() {
let (s, body) = parse(b"#!lua version=5.3").unwrap();
assert_eq!(s.version, LuaVersion::Lua53);
assert_eq!(body, b"");
}
}