1use std::{collections::HashMap, convert::Infallible, fmt::Display, str::FromStr};
2
3use crate::{
4 lockfile::{OptState, PinnedState},
5 lua_rockspec::{
6 ExternalDependencySpec, PartialOverride, PerPlatform, PlatformOverridable, RockSourceSpec,
7 },
8 package::{PackageName, PackageReq, PackageReqParseError, PackageSpec, PackageVersionReq},
9};
10use miette::Diagnostic;
11use serde::{Deserialize, Deserializer};
12use thiserror::Error;
13
14#[derive(Error, Debug, Diagnostic)]
15#[non_exhaustive]
16pub enum LuaDependencySpecParseError {
17 #[error(transparent)]
18 #[diagnostic(transparent)]
19 PackageReq(#[from] PackageReqParseError),
20}
21
22#[derive(Debug, Clone, PartialEq)]
24pub struct LuaDependencySpec {
25 pub(crate) package_req: PackageReq,
26 pub(crate) pin: PinnedState,
27 pub(crate) opt: OptState,
28 pub(crate) source: Option<RockSourceSpec>,
29}
30
31impl LuaDependencySpec {
32 pub fn package_req(&self) -> &PackageReq {
33 &self.package_req
34 }
35 pub fn pin(&self) -> &PinnedState {
36 &self.pin
37 }
38 pub fn opt(&self) -> &OptState {
39 &self.opt
40 }
41 pub fn source(&self) -> &Option<RockSourceSpec> {
42 &self.source
43 }
44 pub fn into_package_req(self) -> PackageReq {
45 self.package_req
46 }
47 pub fn name(&self) -> &PackageName {
48 self.package_req.name()
49 }
50 pub fn version_req(&self) -> &PackageVersionReq {
51 self.package_req.version_req()
52 }
53 pub fn matches(&self, package: &PackageSpec) -> bool {
54 self.package_req.matches(package)
55 }
56}
57
58impl From<PackageName> for LuaDependencySpec {
59 fn from(name: PackageName) -> Self {
60 Self {
61 package_req: PackageReq::from(name),
62 pin: PinnedState::default(),
63 opt: OptState::default(),
64 source: None,
65 }
66 }
67}
68
69impl From<PackageReq> for LuaDependencySpec {
70 fn from(package_req: PackageReq) -> Self {
71 Self {
72 package_req,
73 pin: PinnedState::default(),
74 opt: OptState::default(),
75 source: None,
76 }
77 }
78}
79
80impl FromStr for LuaDependencySpec {
81 type Err = LuaDependencySpecParseError;
82
83 fn from_str(str: &str) -> Result<Self, LuaDependencySpecParseError> {
84 let package_req = PackageReq::from_str(str)?;
85 Ok(Self {
86 package_req,
87 pin: PinnedState::default(),
88 opt: OptState::default(),
89 source: None,
90 })
91 }
92}
93
94impl Display for LuaDependencySpec {
95 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
96 if self.version_req().is_any() {
97 self.name().fmt(f)
98 } else {
99 f.write_str(format!("{}{}", self.name(), self.version_req()).as_str())
100 }
101 }
102}
103
104impl PartialOverride for Vec<LuaDependencySpec> {
108 type Err = Infallible;
109
110 fn apply_overrides(&self, override_vec: &Self) -> Result<Self, Self::Err> {
111 let mut result_map: HashMap<String, LuaDependencySpec> = self
112 .iter()
113 .map(|dep| (dep.name().clone().to_string(), dep.clone()))
114 .collect();
115 for override_dep in override_vec {
116 result_map.insert(
117 override_dep.name().clone().to_string(),
118 override_dep.clone(),
119 );
120 }
121 Ok(result_map.into_values().collect())
122 }
123}
124
125impl PlatformOverridable for Vec<LuaDependencySpec> {
126 type Err = Infallible;
127
128 fn on_nil<T>() -> Result<super::PerPlatform<T>, <Self as PlatformOverridable>::Err>
129 where
130 T: PlatformOverridable,
131 T: Default,
132 {
133 Ok(PerPlatform::default())
134 }
135}
136
137impl<'de> Deserialize<'de> for LuaDependencySpec {
138 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
139 where
140 D: Deserializer<'de>,
141 {
142 let package_req = PackageReq::deserialize(deserializer)?;
143 Ok(Self {
144 package_req,
145 pin: PinnedState::default(),
146 opt: OptState::default(),
147 source: None,
148 })
149 }
150}
151
152#[derive(Debug, Deserialize)]
153#[serde(rename_all = "lowercase")]
154pub enum DependencyType<T> {
155 Regular(Vec<T>),
156 Build(Vec<T>),
157 Test(Vec<T>),
158 External(HashMap<String, ExternalDependencySpec>),
159}
160
161impl<T> DependencyType<T> {
162 pub fn as_ref(&self) -> DependencyType<&T> {
163 match *self {
164 Self::Regular(ref x) => DependencyType::Regular(x.iter().collect()),
165 Self::Build(ref x) => DependencyType::Build(x.iter().collect()),
166 Self::Test(ref x) => DependencyType::Test(x.iter().collect()),
167 Self::External(ref x) => DependencyType::External(x.clone()),
168 }
169 }
170}
171
172#[derive(Debug, Deserialize)]
173#[serde(rename_all = "lowercase")]
174pub enum LuaDependencyType<T> {
175 Regular(Vec<T>),
176 Build(Vec<T>),
177 Test(Vec<T>),
178}
179
180#[cfg(test)]
181mod tests {
182
183 use ottavino::{Closure, Executor, Fuel, Lua, Value};
184 use ottavino_util::serde::from_value;
185 use path_slash::PathBufExt;
186
187 use super::*;
188
189 fn eval_lua<T: serde::de::DeserializeOwned>(code: &str) -> Result<T, ottavino::ExternError> {
190 Lua::core().try_enter(|ctx| {
191 let closure = Closure::load(ctx, None, code.as_bytes())?;
192 let executor = Executor::start(ctx, closure.into(), ());
193 executor.step(ctx, &mut Fuel::with(i32::MAX))?;
194 from_value(executor.take_result::<Value<'_>>(ctx)??).map_err(ottavino::Error::from)
195 })
196 }
197
198 #[tokio::test]
199 async fn test_override_lua_dependency_spec() {
200 let neorg_a: LuaDependencySpec = "neorg 1.0.0".parse().unwrap();
201 let neorg_b: LuaDependencySpec = "neorg 2.0.0".parse().unwrap();
202 let foo: LuaDependencySpec = "foo 1.0.0".parse().unwrap();
203 let bar: LuaDependencySpec = "bar 1.0.0".parse().unwrap();
204 let base_vec = vec![neorg_a, foo.clone()];
205 let override_vec = vec![neorg_b.clone(), bar.clone()];
206 let result = base_vec.apply_overrides(&override_vec).unwrap();
207 assert_eq!(result.clone().len(), 3);
208 assert_eq!(
209 result
210 .into_iter()
211 .filter(|dep| *dep == neorg_b || *dep == foo || *dep == bar)
212 .count(),
213 3
214 );
215 }
216
217 #[test]
218 fn test_dependency_type_from_lua() {
219 let regular_deps: DependencyType<LuaDependencySpec> =
220 eval_lua(r#"return { regular = {"neorg 1.0.0", "foo 1.0.0"} }"#).unwrap();
221 let build_deps: DependencyType<LuaDependencySpec> =
222 eval_lua(r#"return { build = {"neorg 1.0.0", "foo 1.0.0"} }"#).unwrap();
223 let test_deps: DependencyType<LuaDependencySpec> =
224 eval_lua(r#"return { test = {"neorg 1.0.0", "foo 1.0.0"} }"#).unwrap();
225 let external_deps: DependencyType<ExternalDependencySpec> = eval_lua(
226 r#"return { external = { foo = { header = "foo.h", library = "libfoo.so" }, bar = { header = "bar.h" } } }"#,
227 )
228 .unwrap();
229
230 match regular_deps {
231 DependencyType::Regular(deps) => {
232 assert_eq!(deps.len(), 2);
233 assert_eq!(deps[0].to_string(), "neorg==1.0.0");
234 assert_eq!(deps[1].to_string(), "foo==1.0.0");
235 }
236 _ => panic!("Expected regular dependencies"),
237 }
238
239 match build_deps {
240 DependencyType::Build(deps) => {
241 assert_eq!(deps.len(), 2);
242 assert_eq!(deps[0].to_string(), "neorg==1.0.0");
243 assert_eq!(deps[1].to_string(), "foo==1.0.0");
244 }
245 _ => panic!("Expected build dependencies"),
246 }
247
248 match test_deps {
249 DependencyType::Test(deps) => {
250 assert_eq!(deps.len(), 2);
251 assert_eq!(deps[0].to_string(), "neorg==1.0.0");
252 assert_eq!(deps[1].to_string(), "foo==1.0.0");
253 }
254 _ => panic!("Expected test dependencies"),
255 }
256
257 match external_deps {
258 DependencyType::External(deps) => {
259 assert_eq!(deps.len(), 2);
260 assert_eq!(
261 deps["foo"].header.as_ref().unwrap().to_slash_lossy(),
262 "foo.h"
263 );
264 assert_eq!(
265 deps["foo"].library.as_ref().unwrap().to_slash_lossy(),
266 "libfoo.so"
267 );
268
269 assert_eq!(
270 deps["bar"].header.as_ref().unwrap().to_slash_lossy(),
271 "bar.h"
272 );
273 assert!(deps["bar"].library.is_none());
274 }
275 _ => panic!("Expected external dependencies"),
276 }
277
278 let _err: ottavino::ExternError =
279 eval_lua::<DependencyType<ExternalDependencySpec>>("return {}").unwrap_err();
280 }
281
282 #[test]
283 fn test_lua_dependency_type_from_lua() {
284 let regular_deps: LuaDependencyType<LuaDependencySpec> =
285 eval_lua(r#"return { regular = {"neorg 1.0.0", "foo 1.0.0"} }"#).unwrap();
286 let build_deps: LuaDependencyType<LuaDependencySpec> =
287 eval_lua(r#"return { build = {"neorg 1.0.0", "foo 1.0.0"} }"#).unwrap();
288 let test_deps: LuaDependencyType<LuaDependencySpec> =
289 eval_lua(r#"return { test = {"neorg 1.0.0", "foo 1.0.0"} }"#).unwrap();
290
291 match regular_deps {
292 LuaDependencyType::Regular(deps) => {
293 assert_eq!(deps.len(), 2);
294 assert_eq!(deps[0].to_string(), "neorg==1.0.0");
295 assert_eq!(deps[1].to_string(), "foo==1.0.0");
296 }
297 _ => panic!("Expected regular dependencies"),
298 }
299
300 match build_deps {
301 LuaDependencyType::Build(deps) => {
302 assert_eq!(deps.len(), 2);
303 assert_eq!(deps[0].to_string(), "neorg==1.0.0");
304 assert_eq!(deps[1].to_string(), "foo==1.0.0");
305 }
306 _ => panic!("Expected build dependencies"),
307 }
308
309 match test_deps {
310 LuaDependencyType::Test(deps) => {
311 assert_eq!(deps.len(), 2);
312 assert_eq!(deps[0].to_string(), "neorg==1.0.0");
313 assert_eq!(deps[1].to_string(), "foo==1.0.0");
314 }
315 _ => panic!("Expected test dependencies"),
316 }
317
318 eval_lua::<LuaDependencyType<LuaDependencySpec>>("return {}").unwrap_err();
319 }
320}