poolster-plugin-java 0.5.0-alpha.1

Native Java SDK generator for Poolster
Documentation
use anyhow::{Context, Result, bail};
use poolster_core::customization::BundledMiddleware;
use poolster_core::{GeneratedFile, GeneratedTree};

pub(crate) fn bundle(tree: &mut GeneratedTree, middleware: &[BundledMiddleware]) -> Result<()> {
    if middleware.is_empty() {
        return Ok(());
    }
    let middleware = middleware
        .iter()
        .map(|entry| {
            entry.validate()?;
            let mut entry = entry.clone();
            entry.path = entry
                .path
                .components()
                .filter(|part| !matches!(part, std::path::Component::CurDir))
                .collect();
            Ok(entry)
        })
        .collect::<Result<Vec<_>>>()?;
    let middleware = middleware.as_slice();
    let (target, original) = tree
        .iter()
        .find(|(_, text)| text.contains(ANCHOR))
        .map(|(path, text)| (path.to_owned(), text.to_owned()))
        .context("bundled middleware requires the native SDK runtime")?;
    if original.matches(ANCHOR).count() != 1 {
        bail!("unsupported native SDK middleware runtime ABI");
    }
    let source_directory = target
        .parent()
        .unwrap()
        .components()
        .filter(|part| !matches!(part, std::path::Component::CurDir))
        .collect::<std::path::PathBuf>();
    let mut staged = tree.clone();
    for entry in middleware {
        entry.validate()?;
        if entry.async_symbol.is_some() {
            bail!(
                "native middleware does not support async_symbol; wrap the native transport instead"
            );
        }
        if entry.path.extension().and_then(|value| value.to_str()) != Some("java")
            || entry.path.parent() != Some(source_directory.as_path())
            || entry.path.file_stem().and_then(|value| value.to_str())
                != Some(entry.symbol.as_str())
        {
            bail!(
                "Java middleware must be {}/{}.java in the generated client package",
                source_directory.display(),
                entry.symbol
            );
        }
        if staged.iter().any(|(path, _)| {
            path.components()
                .filter(|part| !matches!(part, std::path::Component::CurDir))
                .collect::<std::path::PathBuf>()
                == entry.path
        }) {
            bail!(
                "bundled middleware collides with generated file {}",
                entry.path.display()
            );
        }
        staged.insert(GeneratedFile::new(&entry.path, entry.contents.clone())?)?;
        staged.set_owner(&entry.path, format!("bundled-middleware:{}", entry.symbol))?;
    }
    let mut wrapped = "(config.httpClient() != null ? config.httpClient() : HttpClient.newBuilder().connectTimeout(timeout).build())".to_owned();
    for entry in middleware.iter().rev() {
        wrapped = format!("{}.wrap({wrapped})", entry.symbol);
    }
    let updated = original.replacen(
        ANCHOR,
        &format!("this.httpClient = Objects.requireNonNull({wrapped}, \"middleware transport\");"),
        1,
    );
    staged.replace(GeneratedFile::new(&target, updated)?)?;
    let readme = staged
        .iter()
        .find(|(path, _)| path.file_name().is_some_and(|name| name == "README.md"))
        .map(|(path, contents)| (path.to_owned(), contents.to_owned()));
    if let Some((path, contents)) = readme {
        staged.replace(GeneratedFile::new(
            path,
            format!(
                "{contents}\n## Bundled customer middleware\n\n{}\n",
                DOCUMENTATION
            ),
        )?)?;
    }
    *tree = staged;
    Ok(())
}

const ANCHOR: &str = "this.httpClient = config.httpClient() != null ? config.httpClient() : HttpClient.newBuilder().connectTimeout(timeout).build();";
const DOCUMENTATION: &str = r#"Each configured source is compiled and enabled automatically. Define the configured class in the generated client package, with `public static java.net.http.HttpClient wrap(java.net.http.HttpClient next)`. It must return a transport decorator preserving all abstract HttpClient methods and both sendAsync overloads. The SDK passes its injected client or its default client; callers need no middleware configuration. Declaration order is outermost first. Decorators see each HTTP attempt including SDK retries. Preserve body handler types and interruption; manage streaming body ownership deliberately. No async_symbol or other ABI is supported."#;

#[cfg(test)]
mod tests {
    use super::*;
    fn api() -> poolster_core::Api {
        poolster_core::Api {
            name: "Example".into(),
            version: "1.0.0".into(),
            ..Default::default()
        }
    }
    fn tree() -> GeneratedTree {
        crate::render_sdk_with_policy(
            &api(),
            ".",
            None,
            poolster_core::SdkClientStyle::Namespaced,
            false,
        )
        .unwrap()
    }
    fn entry(path: &std::path::Path, symbol: &str) -> BundledMiddleware {
        BundledMiddleware {
            path: path.to_owned(),
            symbol: symbol.into(),
            contents: "customer-owned source".into(),
            async_symbol: None,
        }
    }
    #[test]
    fn bundles_owned_sources_with_automatic_ordered_registration() {
        let mut generated = tree();
        let target = generated
            .iter()
            .find(|(_, text)| text.contains(ANCHOR))
            .unwrap()
            .0
            .to_owned();
        let first = entry(&target.parent().unwrap().join("First.java"), "First");
        let second = entry(&target.parent().unwrap().join("Second.java"), "Second");
        bundle(&mut generated, &[first.clone(), second.clone()]).unwrap();
        let source = generated.get(&target).unwrap();
        assert!(source.contains("First.wrap(Second.wrap("));
        let normalized = first
            .path
            .components()
            .filter(|part| !matches!(part, std::path::Component::CurDir))
            .collect::<std::path::PathBuf>();
        assert_eq!(generated.get(&normalized), Some("customer-owned source"));
        assert!(generated.iter().any(|(path, text)| {
            path.file_name().is_some_and(|name| name == "README.md")
                && text.contains("Bundled customer middleware")
        }));
        assert!(source.contains("config.httpClient() != null"));
    }
    #[test]
    fn rejects_collisions_and_unsupported_abi_without_partial_changes() {
        let mut generated = tree();
        let target = generated
            .iter()
            .find(|(_, text)| text.contains(ANCHOR))
            .unwrap()
            .0
            .to_owned();
        let first = entry(&target.parent().unwrap().join("First.java"), "First");
        let collision = first.clone();
        let before = generated.clone();
        assert!(
            bundle(&mut generated, &[first.clone(), collision])
                .unwrap_err()
                .to_string()
                .contains("collides")
        );
        assert_eq!(before, generated);
        let mut asynchronous = first;
        asynchronous.async_symbol = Some("Async".into());
        assert!(bundle(&mut generated, &[asynchronous]).is_err());
        assert_eq!(before, generated);
    }
    #[test]
    #[ignore = "requires Maven and JDK 17; compiles bundled customer class and generated registration"]
    fn bundled_middleware_native_compile_and_default_registration() {
        let mut generated = tree();
        let (target, source) = generated
            .iter()
            .find(|(_, text)| text.contains(ANCHOR))
            .map(|(path, text)| (path.to_owned(), text.to_owned()))
            .unwrap();
        let namespace = source
            .lines()
            .find_map(|line| {
                line.strip_prefix("package ")
                    .map(|value| value.trim_end_matches(';').to_owned())
            })
            .unwrap();
        let path = target.parent().unwrap().join("CustomerPolicy.java");
        let mut customer = entry(&path, "CustomerPolicy");
        customer.contents = format!(
            "package {namespace};\npublic final class CustomerPolicy {{ public static int calls = 0; public static java.net.http.HttpClient wrap(java.net.http.HttpClient next) {{ calls++; return next; }} }}\n"
        );
        bundle(&mut generated, &[customer]).unwrap();
        let root = tempfile::tempdir().unwrap();
        generated.write_to(root.path()).unwrap();

        let probe = format!(
            "package {namespace}; public final class MiddlewareProbe {{ public static void main(String[] args) {{ new Client(new ClientConfig(\"https://example.test\", null)); if (CustomerPolicy.calls != 1) throw new AssertionError(\"Not automatically registered\"); }} }}"
        );
        std::fs::write(
            root.path()
                .join(target.parent().unwrap())
                .join("MiddlewareProbe.java"),
            probe,
        )
        .unwrap();
        let output = std::process::Command::new("mvn")
            .args([
                "--batch-mode",
                "--no-transfer-progress",
                "-q",
                "-DskipTests",
                "compile",
                "dependency:build-classpath",
                "-Dmdep.outputFile=dependencies.txt",
            ])
            .current_dir(root.path())
            .output()
            .unwrap();
        assert!(
            output.status.success(),
            "{}{}",
            String::from_utf8_lossy(&output.stdout),
            String::from_utf8_lossy(&output.stderr)
        );
        let dependencies = std::fs::read_to_string(root.path().join("dependencies.txt")).unwrap();
        let classpath = format!(
            "{}:{}",
            root.path().join("target/classes").display(),
            dependencies.trim()
        );
        let output = std::process::Command::new("java")
            .args(["-cp", &classpath, &format!("{namespace}.MiddlewareProbe")])
            .output()
            .unwrap();
        assert!(
            output.status.success(),
            "{}{}",
            String::from_utf8_lossy(&output.stdout),
            String::from_utf8_lossy(&output.stderr)
        );
    }
}