diff --git a/csharp/src/Apache.Arrow.Adbc/C/CAdbcDriverImporter.cs b/csharp/src/Apache.Arrow.Adbc/C/CAdbcDriverImporter.cs index 713fbe4d32..b250665e18 100644 --- a/csharp/src/Apache.Arrow.Adbc/C/CAdbcDriverImporter.cs +++ b/csharp/src/Apache.Arrow.Adbc/C/CAdbcDriverImporter.cs @@ -50,7 +50,20 @@ public static partial class CAdbcDriverImporter /// The path to the driver to load /// Whether the driver can be safely unloaded /// The name of the entry point. If not provided, the name AdbcDriverInit will be used. - public static AdbcDriver Load(string file, bool canUnload, string? entryPoint = null) + public static AdbcDriver Load(string file, bool canUnload, string? entryPoint = null) => + Load(file, canUnload, entryPoint, fallbackEntryPoint: null); + + /// + /// Loads a native driver, trying the fallback entry point when the requested symbol is absent. + /// + internal static AdbcDriver LoadWithFallback(string file, string entryPoint, string fallbackEntryPoint) => + Load(file, canUnload: false, entryPoint, fallbackEntryPoint); + + private static AdbcDriver Load( + string file, + bool canUnload, + string? entryPoint, + string? fallbackEntryPoint) { if (file == null) { @@ -71,7 +84,15 @@ public static AdbcDriver Load(string file, bool canUnload, string? entryPoint = try { entryPoint = entryPoint ?? driverInit; - IntPtr export = NativeLibrary.GetExport(library, entryPoint); + IntPtr export; + try + { + export = NativeLibrary.GetExport(library, entryPoint); + } + catch (EntryPointNotFoundException) when (fallbackEntryPoint != null) + { + export = NativeLibrary.GetExport(library, fallbackEntryPoint); + } if (export == IntPtr.Zero) { if (canUnload) { NativeLibrary.Free(library); } diff --git a/csharp/src/Apache.Arrow.Adbc/DriverManager/AdbcDriverManager.cs b/csharp/src/Apache.Arrow.Adbc/DriverManager/AdbcDriverManager.cs index 3badc3e457..dfae3f468f 100644 --- a/csharp/src/Apache.Arrow.Adbc/DriverManager/AdbcDriverManager.cs +++ b/csharp/src/Apache.Arrow.Adbc/DriverManager/AdbcDriverManager.cs @@ -36,6 +36,8 @@ namespace Apache.Arrow.Adbc.DriverManager /// public static class AdbcDriverManager { + private const string DefaultNativeEntrypoint = "AdbcDriverInit"; + /// /// The environment variable that specifies additional driver search paths. /// @@ -132,7 +134,12 @@ private static AdbcDriver LoadNativeDriver(string driverPath, string? entrypoint typeName: null, manifestPath: null, loadMethod: loadMethod, - () => CAdbcDriverImporter.Load(driverPath, resolvedEntrypoint)); + () => entrypoint == null + ? CAdbcDriverImporter.LoadWithFallback( + driverPath, + resolvedEntrypoint, + DefaultNativeEntrypoint) + : CAdbcDriverImporter.Load(driverPath, resolvedEntrypoint)); } // ----------------------------------------------------------------------- @@ -641,7 +648,7 @@ public static string DeriveEntrypoint(string driverPath) baseName = baseName.Substring(adbcPrefix.Length); if (string.IsNullOrEmpty(baseName)) - return "AdbcDriverInit"; + return DefaultNativeEntrypoint; // Convert snake_case to PascalCase. string pascal = ToPascalCase(baseName); @@ -718,7 +725,10 @@ private static AdbcDriver LoadFromManifest(string manifestPath, string? entrypoi DriverManifest manifest = DriverManifest.LoadFromFile(manifestPath); // Caller-supplied entrypoint wins over the manifest's. Falls back to - // a derived native symbol name only when neither is provided. + // a derived native symbol name only when neither is provided. Keep + // that provenance so the standard native fallback is not applied to + // an explicit caller- or manifest-selected symbol. + bool allowNativeFallback = entrypoint == null && manifest.Entrypoint == null; string resolvedEntrypoint = entrypoint ?? manifest.Entrypoint ?? DeriveEntrypoint(manifest.LibraryPath); @@ -726,7 +736,12 @@ private static AdbcDriver LoadFromManifest(string manifestPath, string? entrypoi string? manifestDir = Path.GetDirectoryName(Path.GetFullPath(manifestPath)); string resolvedPath = ResolveManifestPath(manifest.LibraryPath, manifestDir); - return LoadByEntrypointScheme(resolvedPath, resolvedEntrypoint, manifestPath, nameof(LoadFromManifest)); + return LoadByEntrypointScheme( + resolvedPath, + resolvedEntrypoint, + manifestPath, + nameof(LoadFromManifest), + allowNativeFallback); } /// Returns true if uses a managed-runtime scheme prefix. @@ -763,7 +778,8 @@ private static AdbcDriver LoadByEntrypointScheme( string driverPath, string entrypoint, string? manifestPath, - string loadMethod) + string loadMethod, + bool allowNativeFallback = false) { if (entrypoint.StartsWith(DotnetEntrypointScheme, StringComparison.Ordinal)) { @@ -794,7 +810,12 @@ private static AdbcDriver LoadByEntrypointScheme( typeName: null, manifestPath: manifestPath, loadMethod: loadMethod, - () => CAdbcDriverImporter.Load(driverPath, entrypoint)); + () => allowNativeFallback + ? CAdbcDriverImporter.LoadWithFallback( + driverPath, + entrypoint, + DefaultNativeEntrypoint) + : CAdbcDriverImporter.Load(driverPath, entrypoint)); } /// diff --git a/csharp/test/AotInterop/Apache.Arrow.Adbc.TestFixture.Tests/AotFixtureTests.cs b/csharp/test/AotInterop/Apache.Arrow.Adbc.TestFixture.Tests/AotFixtureTests.cs index f7042b6b64..ccbc15a477 100644 --- a/csharp/test/AotInterop/Apache.Arrow.Adbc.TestFixture.Tests/AotFixtureTests.cs +++ b/csharp/test/AotInterop/Apache.Arrow.Adbc.TestFixture.Tests/AotFixtureTests.cs @@ -50,10 +50,15 @@ public class AotFixtureTests private static AdbcDriver LoadFixture() { string? path = ResolveFixturePath(); + SkipIfFixtureUnavailable(path); + return CAdbcDriverImporter.Load(path!); + } + + private static void SkipIfFixtureUnavailable(string? path) + { Skip.IfNot( path != null, $"Set {FixturePathEnvVar} to the AOT-published Apache.Arrow.Adbc.TestFixture.Native shared library to run this test."); - return CAdbcDriverImporter.Load(path!); } [SkippableFact] @@ -63,6 +68,68 @@ public void DriverNegotiatesV1_1_0() Assert.Equal(AdbcVersion.Version_1_1_0, driver.DriverVersion); } + [SkippableFact] + public void DriverManagerFallsBackToStandardEntrypoint() + { + string? fixturePath = ResolveFixturePath(); + SkipIfFixtureUnavailable(fixturePath); + string? fixtureParent = Path.GetDirectoryName(fixturePath); + Skip.IfNot(!string.IsNullOrEmpty(fixtureParent), "The AOT fixture path must have a parent directory."); + string directory = Path.Combine(fixtureParent!, Guid.NewGuid().ToString("N")); + Directory.CreateDirectory(directory); + string renamedPath = Path.Combine(directory, "adbc_driver_fixture" + Path.GetExtension(fixturePath)); + File.Copy(fixturePath!, renamedPath); + + try + { + using AdbcDriver driver = Apache.Arrow.Adbc.DriverManager.AdbcDriverManager.LoadDriver(renamedPath); + Assert.Equal(AdbcVersion.Version_1_1_0, driver.DriverVersion); + } + finally + { + // DriverManager keeps native drivers loaded for the process lifetime. + // Windows therefore prevents deleting the copied DLL while the test + // process is alive; the CI workspace is disposable. + if (!OperatingSystem.IsWindows()) + { + Directory.Delete(directory, recursive: true); + } + } + } + + [SkippableFact] + public void FindLoadDriverManifestWithoutEntrypointFallsBackToStandardEntrypoint() + { + string? fixturePath = ResolveFixturePath(); + SkipIfFixtureUnavailable(fixturePath); + string? fixtureParent = Path.GetDirectoryName(fixturePath); + Skip.IfNot(!string.IsNullOrEmpty(fixtureParent), "The AOT fixture path must have a parent directory."); + string directory = Path.Combine(fixtureParent!, Guid.NewGuid().ToString("N")); + Directory.CreateDirectory(directory); + string renamedName = "adbc_driver_fixture"; + string renamedPath = Path.Combine(directory, renamedName + Path.GetExtension(fixturePath)); + File.Copy(fixturePath!, renamedPath); + File.WriteAllText( + Path.Combine(directory, renamedName + ".toml"), + "manifest_version = 1\n[Driver]\nshared = \"" + Path.GetFileName(renamedPath) + "\"\n"); + + try + { + using AdbcDriver driver = Apache.Arrow.Adbc.DriverManager.AdbcDriverManager.FindLoadDriver( + renamedName, + loadOptions: Apache.Arrow.Adbc.DriverManager.AdbcLoadFlags.Default, + additionalSearchPathList: directory); + Assert.Equal(AdbcVersion.Version_1_1_0, driver.DriverVersion); + } + finally + { + if (!OperatingSystem.IsWindows()) + { + Directory.Delete(directory, recursive: true); + } + } + } + [SkippableFact] public async Task ExecuteQueryRoundTripsThroughAotBoundary() {