Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
245 changes: 245 additions & 0 deletions src/cloud-sql-instance.ts
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
// See the License for the specific language governing permissions and
// limitations under the License.

import tls from 'node:tls';
import {IpAddressTypes, selectIpAddress} from './ip-addresses';
import {InstanceConnectionInfo} from './instance-connection-info';
import {
Expand All @@ -28,6 +29,10 @@ import {SslCert} from './ssl-cert';
import {getRefreshInterval, isExpirationTimeValid} from './time';
import {AuthTypes} from './auth-types';
import {CloudSQLConnectorError} from './errors';
import {validateCertificate} from './socket';

export const DEFAULT_SERVER_PROXY_PORT = 3307;
export const DEFAULT_CONNECT_TIMEOUT_MS = 30 * 1000;

// Private types that describe exactly the methods
// needed from tls.Socket to be able to close
Expand Down Expand Up @@ -100,6 +105,7 @@ export class CloudSQLInstance {
private closed = false;
private failoverPeriod: number;
private sockets = new Set<DestroyableSocket>();
private iamPrincipals = new Map<string, {user: string; database: string}>();

public readonly instanceInfo: InstanceConnectionInfo;
public ephemeralCert?: SslCert;
Expand Down Expand Up @@ -340,6 +346,10 @@ export class CloudSQLInstance {
serverCaCert,
};

if (this.authType === AuthTypes.IAM) {
await this.probeConnection(nextValues, metadata);
}

// In the rather odd case that the current ephemeral certificate is still
// valid while we get an invalid result from the API calls, then preserve
// the current metadata.
Expand All @@ -350,6 +360,122 @@ export class CloudSQLInstance {
return nextValues;
}

recordIamPrincipal(user: string, database: string): void {
if (!user) {
return;
}
const db = database || user;
this.iamPrincipals.set(`${user}\0${db}`, {user, database: db});
}

private async probeConnection(
refreshResult: RefreshResult,
metadata: InstanceMetadata
): Promise<void> {
const {ephemeralCert, privateKey, serverCaCert} = refreshResult;
if (!ephemeralCert || !privateKey || !serverCaCert) {
return;
}

const targets: string[] = [];
if (this.instanceInfo && this.instanceInfo.domainName) {
targets.push(this.instanceInfo.domainName);
} else {
try {
const selectedIp = selectIpAddress(metadata.ipAddresses, this.ipType);
if (selectedIp) {
targets.push(selectedIp);
}
} catch {
// If the configured IP type is not available in metadata, skip probe
}
}

if (targets.length === 0) {
return;
}

const principals: Array<{user: string; database: string} | null> =
this.iamPrincipals.size > 0
? Array.from(this.iamPrincipals.values())
: [null];

const port = this.port || DEFAULT_SERVER_PROXY_PORT;
for (const principal of principals) {
for (const target of targets) {
try {
await new Promise<void>((resolve, reject) => {
let settled = false;
const finish = (err?: Error) => {
if (settled) {
return;
}
settled = true;
clearTimeout(timeout);
if (err) {
reject(err);
} else {
resolve();
}
};

const timeout = setTimeout(() => {
socket.destroy(new Error('Probe timeout'));
finish(new Error('Probe timeout'));
}, DEFAULT_CONNECT_TIMEOUT_MS);

const socket: tls.TLSSocket = tls.connect(
{
host: target,
port,
secureContext: tls.createSecureContext({
ca: serverCaCert.cert,
cert: ephemeralCert.cert,
key: privateKey,
minVersion: 'TLSv1.3',
}),
checkServerIdentity: validateCertificate(
this.instanceInfo,
metadata.dnsName || '',
target
),
},
() => {
if (!principal) {
socket.end();
finish();
return;
}
socket.once('data', () => {
socket.write(Buffer.from([0x58, 0x00, 0x00, 0x00, 0x04]));
socket.end();
finish();
});
socket.once('end', () => {
finish();
});
socket.once('close', () => {
finish();
});
socket.write(
buildPostgresStartupPacket(principal.user, principal.database)
);
}
);

socket.on('error', err => {
socket.destroy();
finish(err);
});
});
break;
} catch (e) {
// Ignore probe error across single target and try next target
}
}
}
}

private isValid({
ephemeralCert,
host,
Expand Down Expand Up @@ -465,6 +591,9 @@ export class CloudSQLInstance {
return false;
}
addSocket(socket: DestroyableSocket) {
if (this.authType === AuthTypes.IAM) {
this.attachPostgresStartupSniffer(socket);
}
if (!this.instanceInfo.domainName) {
// This was not connected by domain name. Ignore all sockets.
return;
Expand All @@ -477,4 +606,120 @@ export class CloudSQLInstance {
this.sockets.delete(socket);
});
}

private attachPostgresStartupSniffer(socket: DestroyableSocket): void {
const writable = socket as unknown as {
write?: (...args: unknown[]) => boolean;
};
if (typeof writable.write !== 'function') {
return;
}
const origWrite = writable.write;
let buf: Buffer = Buffer.alloc(0);
let done = false;
writable.write = (...args: unknown[]): boolean => {
if (done) {
return origWrite.apply(socket, args);
}
const chunk = args[0];
const chunkBuf: Buffer | null = Buffer.isBuffer(chunk)
? chunk
: typeof chunk === 'string'
? Buffer.from(chunk)
: chunk instanceof Uint8Array
? Buffer.from(chunk)
: null;
if (chunkBuf) {
buf = Buffer.concat([buf, chunkBuf]);
const parsed = parsePostgresStartupPacket(buf);
if (parsed.complete) {
done = true;
buf = Buffer.alloc(0);
writable.write = origWrite;
if (parsed.user) {
this.recordIamPrincipal(parsed.user, parsed.database);
}
} else if (buf.length > MAX_PG_STARTUP_PACKET_LEN + 8) {
done = true;
buf = Buffer.alloc(0);
writable.write = origWrite;
}
}
return origWrite.apply(socket, args);
};
}
}

const PG_SSL_REQUEST_CODE = 80877103; // 0x04d2162f
const PG_PROTOCOL_VERSION_30 = 196608; // 0x00030000
const MAX_PG_STARTUP_PACKET_LEN = 10000;

export function parsePostgresStartupPacket(buf: Buffer): {
user: string;
database: string;
complete: boolean;
} {
let slice = buf;
if (slice.length < 8) {
return {user: '', database: '', complete: false};
}
let pktLen = slice.readUInt32BE(0);
let code = slice.readUInt32BE(4);
if (pktLen === 8 && code === PG_SSL_REQUEST_CODE) {
slice = slice.subarray(8);
if (slice.length < 8) {
return {user: '', database: '', complete: false};
}
pktLen = slice.readUInt32BE(0);
code = slice.readUInt32BE(4);
}
if (
code !== PG_PROTOCOL_VERSION_30 ||
pktLen < 8 ||
pktLen > MAX_PG_STARTUP_PACKET_LEN
) {
return {user: '', database: '', complete: true};
}
if (slice.length < pktLen) {
return {user: '', database: '', complete: false};
}
let payload = slice.subarray(8, pktLen);
let user = '';
let database = '';
while (payload.length > 0 && payload[0] !== 0) {
const kEnd = payload.indexOf(0);
if (kEnd < 0) {
break;
}
const key = payload.subarray(0, kEnd).toString('utf8');
payload = payload.subarray(kEnd + 1);
const vEnd = payload.indexOf(0);
if (vEnd < 0) {
break;
}
const val = payload.subarray(0, vEnd).toString('utf8');
payload = payload.subarray(vEnd + 1);
if (key === 'user') {
user = val;
} else if (key === 'database') {
database = val;
}
}
if (user && !database) {
database = user;
}
return {user, database, complete: true};
}

export function buildPostgresStartupPacket(
user: string,
database: string
): Buffer {
const db = database || user;
const body = Buffer.from(`user\0${user}\0database\0${db}\0\0`, 'utf8');
const totalLen = 8 + body.length;
const header = Buffer.alloc(8);
header.writeUInt32BE(totalLen, 0);
header.writeUInt32BE(PG_PROTOCOL_VERSION_30, 4);
return Buffer.concat([header, body]);
}
91 changes: 91 additions & 0 deletions test/cloud-sql-instance.ts
Original file line number Diff line number Diff line change
Expand Up @@ -703,4 +703,95 @@ t.test('CloudSQLInstance', async t => {
'socket not added when domainName not set'
);
});

t.test(
'probeConnection sends PostgreSQL v3 StartupMessage and Terminate when IAM principal is recorded',
async t => {
const {EventEmitter} = await import('node:events');
const writtenPackets: Buffer[] = [];
let ended = false;

const {CloudSQLInstance: MockedInstance, buildPostgresStartupPacket} =
t.mockRequire('../src/cloud-sql-instance', {
'../src/crypto': {
generateKeys: async () => ({
publicKey: '-----BEGIN PUBLIC KEY-----',
privateKey: CLIENT_KEY,
}),
},
'../src/time': {
getRefreshInterval() {
return 50;
},
isExpirationTimeValid() {
return true;
},
},
'node:tls': {
createSecureContext: () => ({}),
connect: (_opts: unknown, onSecureConnect: () => void) => {
const fakeSocket = new EventEmitter() as EventEmitter & {
write: (buf: Buffer) => boolean;
end: () => void;
destroy: () => void;
};
fakeSocket.write = (buf: Buffer) => {
writtenPackets.push(Buffer.from(buf));
if (writtenPackets.length === 1) {
setImmediate(() => {
fakeSocket.emit(
'data',
Buffer.from([0x52, 0, 0, 0, 8, 0, 0, 0, 0])
);
});
}
return true;
};
fakeSocket.end = () => {
ended = true;
};
fakeSocket.destroy = () => {};
setImmediate(onSecureConnect);
return fakeSocket;
},
},
});

const instance = new MockedInstance({
options: {
ipType: IpAddressTypes.PUBLIC,
authType: AuthTypes.IAM,
instanceConnectionName: 'my-project:us-east1:my-instance',
sqlAdminFetcher: fetcher,
limitRateInterval: 50,
},
});
t.after(() => instance.close());

// Simulate an application socket capturing (user, database) via addSocket
const clientSocket = {
write: (chunk: Buffer) => chunk.length >= 0,
destroy: () => {},
once: () => {},
};
instance.addSocket(clientSocket);
const sslReq = Buffer.from([0, 0, 0, 8, 0x04, 0xd2, 0x16, 0x2f]);
clientSocket.write(sslReq);
const startup = buildPostgresStartupPacket(
'iam-user@example.com',
'mydb'
);
clientSocket.write(startup);

await instance.refresh();
instance.cancelRefresh();

t.same(
writtenPackets,
[startup, Buffer.from([0x58, 0x00, 0x00, 0x00, 0x04])],
'should write PostgreSQL StartupMessage and Terminate on probe'
);
t.ok(ended, 'should end probe socket after Terminate');
}
);
});
Loading