diff --git a/Npgsql.sln.DotSettings b/Npgsql.sln.DotSettings index 09e69a826c..e4904b88ba 100644 --- a/Npgsql.sln.DotSettings +++ b/Npgsql.sln.DotSettings @@ -91,6 +91,7 @@ True True True + True True True True diff --git a/src/Npgsql/NpgsqlConnector.Auth.cs b/src/Npgsql/NpgsqlConnector.Auth.cs index 8fe527a5aa..bc6d333973 100644 --- a/src/Npgsql/NpgsqlConnector.Auth.cs +++ b/src/Npgsql/NpgsqlConnector.Auth.cs @@ -460,12 +460,7 @@ class AuthenticationCompleteException : Exception { } if (password != null) return password; - var passFile = Settings.Passfile ?? PostgresEnvironment.PassFile; - if (passFile is null && PostgresEnvironment.PassFileDefault is string passFileDefault) - { - passFile = passFileDefault; - } - + var passFile = Settings.Passfile ?? PostgresEnvironment.PassFile ?? PostgresEnvironment.PassFileDefault; if (passFile != null) { var matchingEntry = new PgPassFile(passFile!) diff --git a/src/Npgsql/NpgsqlConnector.cs b/src/Npgsql/NpgsqlConnector.cs index 8f5ae05658..0ad8c44fc6 100644 --- a/src/Npgsql/NpgsqlConnector.cs +++ b/src/Npgsql/NpgsqlConnector.cs @@ -644,7 +644,7 @@ string GetUsername() if (username?.Length > 0) return username; - if (!PGUtil.IsWindows) + if (!RuntimeInformation.IsOSPlatform(OSPlatform.Windows)) { username = KerberosUsernameProvider.GetUsername(Settings.IncludeRealm); if (username?.Length > 0) @@ -716,7 +716,7 @@ async Task RawOpen(NpgsqlTimeout timeout, bool async, CancellationToken cancella { #if NET5_0 // It's PEM time - var keyPath = Settings.SslKey ?? PostgresEnvironment.SslKey; + var keyPath = Settings.SslKey ?? PostgresEnvironment.SslKey ?? PostgresEnvironment.SslKeyDefault; cert = string.IsNullOrEmpty(password) ? X509Certificate2.CreateFromPemFile(certPath, keyPath) : X509Certificate2.CreateFromEncryptedPemFile(certPath, password, keyPath); diff --git a/src/Npgsql/PostgresEnvironment.cs b/src/Npgsql/PostgresEnvironment.cs index 0a664af92a..cc8a0c2a72 100644 --- a/src/Npgsql/PostgresEnvironment.cs +++ b/src/Npgsql/PostgresEnvironment.cs @@ -1,6 +1,7 @@ using System; using System.Collections.Generic; using System.IO; +using System.Runtime.InteropServices; using Npgsql.Util; namespace Npgsql @@ -13,17 +14,28 @@ static class PostgresEnvironment internal static string? PassFile => Environment.GetEnvironmentVariable("PGPASSFILE"); - internal static string? PassFileDefault => GetDefaultFilePath(PGUtil.IsWindows ? "pgpass.conf" : ".pgpass"); + internal static string? PassFileDefault + => (RuntimeInformation.IsOSPlatform(OSPlatform.Windows) + ? Path.Combine(GetHomePostgresDir(), "pgpass.conf") + : Path.Combine(GetHomeDir(), ".pgpass")) is var path && + File.Exists(path) + ? path + : null; internal static string? SslCert => Environment.GetEnvironmentVariable("PGSSLCERT"); - internal static string? SslCertDefault => GetDefaultFilePath("postgresql.crt"); + internal static string? SslCertDefault + => Path.Combine(GetHomePostgresDir(), "postgresql.crt") is var path && File.Exists(path) ? path : null; - internal static string? SslCertRoot => Environment.GetEnvironmentVariable("PGSSLROOTCERT"); + internal static string? SslKey => Environment.GetEnvironmentVariable("PGSSLKEY"); - internal static string? SslCertRootDefault => GetDefaultFilePath("root.crt"); + internal static string? SslKeyDefault + => Path.Combine(GetHomePostgresDir(), "postgresql.key") is var path && File.Exists(path) ? path : null; - internal static string? SslKey => Environment.GetEnvironmentVariable("PGSSLKEY"); + internal static string? SslCertRoot => Environment.GetEnvironmentVariable("PGSSLROOTCERT"); + + internal static string? SslCertRootDefault + => Path.Combine(GetHomePostgresDir(), "root.crt") is var path && File.Exists(path) ? path : null; internal static string? ClientEncoding => Environment.GetEnvironmentVariable("PGCLIENTENCODING"); @@ -33,11 +45,15 @@ static class PostgresEnvironment internal static string? TargetSessionAttributes => Environment.GetEnvironmentVariable("PGTARGETSESSIONATTRS"); - static string? GetDefaultFilePath(string fileName) => - Environment.GetEnvironmentVariable(PGUtil.IsWindows ? "APPDATA" : "HOME") is string appData && - Path.Combine(appData, "postgresql", fileName) is string filePath && - File.Exists(filePath) - ? filePath - : null; + static string GetHomeDir() + { + var envVar = RuntimeInformation.IsOSPlatform(OSPlatform.Windows) ? "APPDATA" : "HOME"; + return Environment.GetEnvironmentVariable(envVar) is string homedir + ? homedir + : throw new InvalidOperationException($"Environment variable {envVar} not defined"); + } + + static string GetHomePostgresDir() + => Path.Combine(GetHomeDir(), RuntimeInformation.IsOSPlatform(OSPlatform.Windows) ? "postgresql" : ".postgresql"); } } diff --git a/src/Npgsql/Util/PGUtil.cs b/src/Npgsql/Util/PGUtil.cs index efcaeea6a5..782f28ef2b 100644 --- a/src/Npgsql/Util/PGUtil.cs +++ b/src/Npgsql/Util/PGUtil.cs @@ -119,9 +119,6 @@ internal static int RotateShift(int val, int shift) internal static readonly Task FalseTask = Task.FromResult(false); internal static StringComparer InvariantCaseIgnoringStringComparer => StringComparer.InvariantCultureIgnoreCase; - - internal static bool IsWindows => - System.Runtime.InteropServices.RuntimeInformation.IsOSPlatform(System.Runtime.InteropServices.OSPlatform.Windows); } enum FormatCode : short diff --git a/test/Npgsql.Tests/ConnectionTests.cs b/test/Npgsql.Tests/ConnectionTests.cs index 2d960386e9..70a3603722 100644 --- a/test/Npgsql.Tests/ConnectionTests.cs +++ b/test/Npgsql.Tests/ConnectionTests.cs @@ -1401,30 +1401,99 @@ public async Task NoResetOnClose(bool noResetOnClose) } [Test] - [NonParallelizable] - public async Task UsePgPassFile() + public async Task Use_pgpass_from_connection_string() { using var resetPassword = SetEnvironmentVariable("PGPASSWORD", null); var builder = new NpgsqlConnectionStringBuilder(ConnectionString); var password = builder.Password; - var passFile = Path.GetTempFileName(); + builder.Password = null; - builder.Password = password; + var passFile = Path.GetTempFileName(); + File.WriteAllText(passFile, $"*:*:*:{builder.Username}:{password}"); builder.Passfile = passFile; - using var deletePassFile = Defer(() => File.Delete(passFile)); + try + { + using var pool = CreateTempPool(builder.ConnectionString, out var connectionString); + using var conn = await OpenConnectionAsync(connectionString); + } + finally + { + File.Delete(passFile); + } + } - File.WriteAllText(passFile, $"*:*:*:{builder.Username}:{password}"); + [Test] + [NonParallelizable] + public async Task Use_pgpass_from_environment_variable() + { + using var resetPassword = SetEnvironmentVariable("PGPASSWORD", null); + var builder = new NpgsqlConnectionStringBuilder(ConnectionString); + var password = builder.Password; + builder.Password = null; + + var passFile = Path.GetTempFileName(); + File.WriteAllText(passFile, $"*:*:*:{builder.Username}:{password}"); using var passFileVariable = SetEnvironmentVariable("PGPASSFILE", passFile); - using var pool = CreateTempPool(builder.ConnectionString, out var connectionString); - using var conn = await OpenConnectionAsync(connectionString); + + try + { + using var pool = CreateTempPool(builder.ConnectionString, out var connectionString); + using var conn = await OpenConnectionAsync(connectionString); + } + finally + { + File.Delete(passFile); + } + } + + [Test] + [NonParallelizable] + public async Task Use_pgpass_from_homedir() + { + using var resetPassword = SetEnvironmentVariable("PGPASSWORD", null); + var builder = new NpgsqlConnectionStringBuilder(ConnectionString); + + var password = builder.Password; + builder.Password = null; + + string? dirToDelete = null; + string passFile; + if (RuntimeInformation.IsOSPlatform(OSPlatform.Windows)) + { + var dir = Path.Combine(Environment.GetEnvironmentVariable("APPDATA")!, "postgresql"); + if (!Directory.Exists(dir)) + { + Directory.CreateDirectory(dir); + dirToDelete = dir; + + } + passFile = Path.Combine(dir, "pgpass.conf"); + } + else + { + passFile = Path.Combine(Environment.GetEnvironmentVariable("HOME")!, ".pgpass"); + } + + try + { + File.WriteAllText(passFile, $"*:*:*:{builder.Username}:{password}"); + using var pool = CreateTempPool(builder.ConnectionString, out var connectionString); + using var conn = await OpenConnectionAsync(connectionString); + } + finally + { + File.Delete(passFile); + if (dirToDelete is not null) + Directory.Delete(dirToDelete); + } } [Test] [NonParallelizable] - public void PasswordSourcePrecendence() + public void PasswordSourcePrecedence() { using var resetPassword = SetEnvironmentVariable("PGPASSWORD", null); var builder = new NpgsqlConnectionStringBuilder(ConnectionString); @@ -1464,7 +1533,7 @@ Func OpenConnection(string? password, string? passFile) => async () = builder.Password = password; builder.Passfile = passFile; builder.IntegratedSecurity = false; - builder.ApplicationName = $"{nameof(PasswordSourcePrecendence)}:{Guid.NewGuid()}"; + builder.ApplicationName = $"{nameof(PasswordSourcePrecedence)}:{Guid.NewGuid()}"; using var pool = CreateTempPool(builder.ConnectionString, out var connectionString); using var connection = await OpenConnectionAsync(connectionString);