Skip to main content

kcode_k1_rust_code_document/
lib.rs

1use kcode_k1_rust_package::{LibraryId, SourceFile, SourcePackage};
2use proc_macro2::Span;
3use std::fmt::{Display, Formatter};
4use std::ops::Range;
5use syn::spanned::Spanned;
6use syn::{ImplItem, Item, TraitItem};
7use toml::{Table, Value};
8
9const SUPPORTED_PATHS: [&str; 3] = ["Cargo.toml", "Documentation.md", "src/lib.rs"];
10const MIN_CHUNK_SCALARS: usize = 800;
11const DEFAULT_DOCUMENTATION: &str = "# Consumer contract\n\n`run` performs no operation.\n";
12const DEFAULT_CODE: &str = "pub fn run() {}\n";
13
14#[derive(Clone, Debug, Eq, PartialEq)]
15pub struct RustCodeDocument {
16    package: SourcePackage,
17}
18
19impl RustCodeDocument {
20    pub fn new(
21        identity: LibraryId,
22        documentation: impl Into<String>,
23        manifest: impl Into<String>,
24        code: impl Into<String>,
25    ) -> Result<Self, RustCodeDocumentError> {
26        let files = vec![
27            SourceFile::new("Documentation.md", documentation.into().into_bytes()),
28            SourceFile::new("Cargo.toml", manifest.into().into_bytes()),
29            SourceFile::new("src/lib.rs", code.into().into_bytes()),
30        ];
31        let package = SourcePackage::new(identity, files)
32            .map_err(|cause| RustCodeDocumentError(cause.to_string()))?;
33        Self::from_source_package(package)
34    }
35
36    pub fn from_source_package(package: SourcePackage) -> Result<Self, RustCodeDocumentError> {
37        let paths: Vec<&str> = package.files().iter().map(|file| file.path()).collect();
38        if paths != SUPPORTED_PATHS {
39            return fail("source package must contain exactly the three supported paths");
40        }
41        for file in package.files() {
42            if std::str::from_utf8(file.bytes()).is_err() {
43                return fail(format!("{} must be UTF-8", file.path()));
44            }
45        }
46        validate_manifest(package.id(), source_text(&package, "Cargo.toml"))?;
47        Ok(Self { package })
48    }
49
50    pub fn approved_default(identity: LibraryId) -> Result<Self, RustCodeDocumentError> {
51        let family = identity.family();
52        let manifest = format!(
53            "[package]\nname = \"{}\"\nversion = \"{}\"\nedition = \"2024\"\nrust-version = \"1.97\"\nlicense = \"MIT\"\nautobins = false\nautoexamples = false\nautotests = false\nautobenches = false\n\n[lib]\nname = \"{}\"\npath = \"src/lib.rs\"\n\n[lints.rust]\nunsafe_code = \"forbid\"\n\n[workspace]\nresolver = \"3\"\n",
54            family.package_name(),
55            identity.version(),
56            family.logical_name().replace('-', "_")
57        );
58        Self::new(identity, DEFAULT_DOCUMENTATION, manifest, DEFAULT_CODE)
59    }
60
61    pub fn identity(&self) -> &LibraryId {
62        self.package.id()
63    }
64
65    pub fn documentation(&self) -> &str {
66        source_text(&self.package, "Documentation.md")
67    }
68
69    pub fn manifest(&self) -> &str {
70        source_text(&self.package, "Cargo.toml")
71    }
72
73    pub fn code(&self) -> &str {
74        source_text(&self.package, "src/lib.rs")
75    }
76
77    pub fn into_source_package(self) -> SourcePackage {
78        self.package
79    }
80
81    pub fn code_chunk_ranges(&self) -> Vec<Range<usize>> {
82        chunk_ranges(self.code())
83    }
84}
85
86#[derive(Clone, Debug, Eq, PartialEq)]
87pub struct RustCodeDocumentError(String);
88
89impl RustCodeDocumentError {
90    pub fn message(&self) -> &str {
91        &self.0
92    }
93}
94
95impl Display for RustCodeDocumentError {
96    fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
97        formatter.write_str(&self.0)
98    }
99}
100
101impl std::error::Error for RustCodeDocumentError {}
102
103fn fail<T>(message: impl Into<String>) -> Result<T, RustCodeDocumentError> {
104    Err(RustCodeDocumentError(message.into()))
105}
106
107fn source_text<'a>(package: &'a SourcePackage, path: &str) -> &'a str {
108    let file = package
109        .files()
110        .iter()
111        .find(|file| file.path() == path)
112        .expect("validated document path");
113    std::str::from_utf8(file.bytes()).expect("validated UTF-8 document")
114}
115
116fn validate_manifest(identity: &LibraryId, source: &str) -> Result<(), RustCodeDocumentError> {
117    let root: Table = source
118        .parse()
119        .map_err(|_| RustCodeDocumentError("Cargo.toml must be valid TOML".into()))?;
120    if root.keys().any(|key| {
121        !matches!(
122            key.as_str(),
123            "package"
124                | "lib"
125                | "dependencies"
126                | "dev-dependencies"
127                | "build-dependencies"
128                | "target"
129                | "lints"
130                | "workspace"
131        )
132    }) {
133        return fail("Cargo.toml contains an unsupported root entry");
134    }
135    let package = required_table(&root, "package")?;
136    let package_keys = [
137        "name",
138        "version",
139        "edition",
140        "rust-version",
141        "license",
142        "autobins",
143        "autoexamples",
144        "autotests",
145        "autobenches",
146    ];
147    if !exact_keys(package, &package_keys) {
148        return fail("[package] must contain exactly the supported entries");
149    }
150    require_text(package, "name", &identity.family().package_name())?;
151    require_text(package, "version", &identity.version().to_string())?;
152    require_text(package, "edition", "2024")?;
153    require_text(package, "rust-version", "1.97")?;
154    require_text(package, "license", "MIT")?;
155    for key in ["autobins", "autoexamples", "autotests", "autobenches"] {
156        if package.get(key) != Some(&Value::Boolean(false)) {
157            return fail(format!("package {key} must be false"));
158        }
159    }
160    let library = required_table(&root, "lib")?;
161    if !exact_keys(library, &["name", "path"]) {
162        return fail("[lib] must contain exactly name and path");
163    }
164    require_text(
165        library,
166        "name",
167        &identity.family().logical_name().replace('-', "_"),
168    )?;
169    require_text(library, "path", "src/lib.rs")?;
170    let lints = required_table(&root, "lints")?;
171    if !exact_keys(lints, &["rust"]) {
172        return fail("[lints] must contain exactly rust");
173    }
174    let rust_lints = required_table(lints, "rust")?;
175    if !exact_keys(rust_lints, &["unsafe_code"])
176        || rust_lints.get("unsafe_code").and_then(Value::as_str) != Some("forbid")
177    {
178        return fail("[lints.rust] must forbid unsafe_code");
179    }
180    let workspace = required_table(&root, "workspace")?;
181    if !exact_keys(workspace, &["resolver"])
182        || workspace.get("resolver").and_then(Value::as_str) != Some("3")
183    {
184        return fail("[workspace] must contain only resolver = 3");
185    }
186    validate_dependency_locations(&root)
187}
188
189fn required_table<'a>(table: &'a Table, key: &str) -> Result<&'a Table, RustCodeDocumentError> {
190    table
191        .get(key)
192        .and_then(Value::as_table)
193        .ok_or_else(|| RustCodeDocumentError(format!("missing or invalid [{key}]")))
194}
195
196fn exact_keys(table: &Table, keys: &[&str]) -> bool {
197    table.len() == keys.len() && keys.iter().all(|key| table.contains_key(*key))
198}
199
200fn require_text(table: &Table, key: &str, expected: &str) -> Result<(), RustCodeDocumentError> {
201    if table.get(key).and_then(Value::as_str) != Some(expected) {
202        return fail(format!("{key} must equal {expected}"));
203    }
204    Ok(())
205}
206
207fn validate_dependency_locations(root: &Table) -> Result<(), RustCodeDocumentError> {
208    let sections = ["dependencies", "dev-dependencies", "build-dependencies"];
209    for section in sections {
210        if root.get(section).is_some_and(|value| !value.is_table()) {
211            return fail(format!("[{section}] must be a table"));
212        }
213    }
214    let Some(targets) = root.get("target") else {
215        return Ok(());
216    };
217    let targets = targets
218        .as_table()
219        .ok_or_else(|| RustCodeDocumentError("[target] must be a table".into()))?;
220    for target in targets.values() {
221        let target = target
222            .as_table()
223            .ok_or_else(|| RustCodeDocumentError("target entry must be a table".into()))?;
224        if target
225            .iter()
226            .any(|(key, value)| !sections.contains(&key.as_str()) || !value.is_table())
227        {
228            return fail("target entries may contain only dependency tables");
229        }
230    }
231    Ok(())
232}
233
234fn whole_range(end: usize) -> Vec<Range<usize>> {
235    std::iter::once(0..end).collect()
236}
237
238fn chunk_ranges(source: &str) -> Vec<Range<usize>> {
239    if source.is_empty() {
240        return whole_range(0);
241    }
242    let Some(points) = safe_points(source) else {
243        return whole_range(source.len());
244    };
245    let mut ranges = Vec::new();
246    let mut start = 0;
247    for end in points {
248        if end <= start || end >= source.len() {
249            continue;
250        }
251        if source[start..end].chars().count() < MIN_CHUNK_SCALARS
252            || source[end..].chars().count() < MIN_CHUNK_SCALARS
253        {
254            continue;
255        }
256        ranges.push(start..end);
257        start = end;
258    }
259    ranges.push(start..source.len());
260    ranges
261}
262
263fn safe_points(source: &str) -> Option<Vec<usize>> {
264    let file = syn::parse_file(source).ok()?;
265    let mut points = Vec::new();
266    for item in &file.items {
267        let outer = checked_span(item.span(), source)?;
268        match item {
269            Item::Impl(item) => {
270                for member in &item.items {
271                    match member {
272                        ImplItem::Const(_)
273                        | ImplItem::Fn(_)
274                        | ImplItem::Type(_)
275                        | ImplItem::Macro(_) => {
276                            add_member(member.span(), &outer, source, &mut points)?
277                        }
278                        ImplItem::Verbatim(_) => return None,
279                        _ => return None,
280                    }
281                }
282            }
283            Item::Trait(item) => {
284                for member in &item.items {
285                    match member {
286                        TraitItem::Const(_)
287                        | TraitItem::Fn(_)
288                        | TraitItem::Type(_)
289                        | TraitItem::Macro(_) => {
290                            add_member(member.span(), &outer, source, &mut points)?
291                        }
292                        TraitItem::Verbatim(_) => return None,
293                        _ => return None,
294                    }
295                }
296            }
297            Item::Const(_)
298            | Item::Enum(_)
299            | Item::ExternCrate(_)
300            | Item::Fn(_)
301            | Item::ForeignMod(_)
302            | Item::Macro(_)
303            | Item::Mod(_)
304            | Item::Static(_)
305            | Item::Struct(_)
306            | Item::TraitAlias(_)
307            | Item::Type(_)
308            | Item::Union(_)
309            | Item::Use(_) => {}
310            Item::Verbatim(_) => return None,
311            _ => return None,
312        }
313        points.push(outer.end);
314    }
315    points.sort_unstable();
316    points.dedup();
317    Some(points)
318}
319
320fn checked_span(span: Span, source: &str) -> Option<Range<usize>> {
321    let range = span.byte_range();
322    if range.start >= range.end
323        || range.end > source.len()
324        || !source.is_char_boundary(range.start)
325        || !source.is_char_boundary(range.end)
326        || source[range.clone()].trim().is_empty()
327    {
328        return None;
329    }
330    Some(range)
331}
332
333fn add_member(
334    span: Span,
335    outer: &Range<usize>,
336    source: &str,
337    points: &mut Vec<usize>,
338) -> Option<()> {
339    let member = checked_span(span, source)?;
340    if member.start < outer.start || member.end > outer.end {
341        return None;
342    }
343    points.push(member.end);
344    Some(())
345}