diff --git a/services/companion/java/com/android/server/companion/securechannel/SecureChannel.java b/services/companion/java/com/android/server/companion/securechannel/SecureChannel.java index 05b6022ce5698..a8519e388525d 100644 --- a/services/companion/java/com/android/server/companion/securechannel/SecureChannel.java +++ b/services/companion/java/com/android/server/companion/securechannel/SecureChannel.java @@ -128,6 +128,9 @@ public class SecureChannel { * Start listening for incoming messages. */ public void start() { + if (DEBUG) { + Slog.d(TAG, "Starting secure channel."); + } new Thread(() -> { try { // 1. Wait for the next handshake message and process it. @@ -151,14 +154,14 @@ public class SecureChannel { // TODO: Handle different types errors. Slog.e(TAG, "Secure channel encountered an error.", e); - stop(); + close(); mCallback.onError(e); } }).start(); } /** - * Stop listening to incoming messages and close the channel. + * Stop listening to incoming messages. */ public void stop() { if (DEBUG) { @@ -166,7 +169,17 @@ public class SecureChannel { } mStopped = true; mInProgress = false; + } + /** + * Stop listening to incoming messages and close the channel. + */ + public void close() { + stop(); + + if (DEBUG) { + Slog.d(TAG, "Closing secure channel."); + } IoUtils.closeQuietly(mInput); IoUtils.closeQuietly(mOutput); KeyStoreUtils.cleanUp(mAlias); @@ -240,60 +253,64 @@ public class SecureChannel { if (isSecured()) { Slog.d(TAG, "Waiting to receive next secure message."); } else { - Slog.d(TAG, "Waiting to receive next message."); + Slog.d(TAG, "Waiting to receive next " + expected + " message."); } } // TODO: Handle message timeout - // Header is _not_ encrypted, but will be covered by MAC - final byte[] headerBytes = new byte[HEADER_LENGTH]; - Streams.readFully(mInput, headerBytes); - final ByteBuffer header = ByteBuffer.wrap(headerBytes); - final int version = header.getInt(); - final short type = header.getShort(); + synchronized (mInput) { + // Header is _not_ encrypted, but will be covered by MAC + final byte[] headerBytes = new byte[HEADER_LENGTH]; + Streams.readFully(mInput, headerBytes); + final ByteBuffer header = ByteBuffer.wrap(headerBytes); + final int version = header.getInt(); + final short type = header.getShort(); - if (version != VERSION) { - Streams.skipByReading(mInput, Long.MAX_VALUE); - throw new SecureChannelException("Secure channel version mismatch. " - + "Currently on version " + VERSION + ". Skipping rest of data."); + if (version != VERSION) { + Streams.skipByReading(mInput, Long.MAX_VALUE); + throw new SecureChannelException("Secure channel version mismatch. " + + "Currently on version " + VERSION + ". Skipping rest of data."); + } + + if (type != expected.mValue) { + Streams.skipByReading(mInput, Long.MAX_VALUE); + throw new SecureChannelException( + "Unexpected message type. Expected " + expected.name() + + "; Found " + MessageType.from(type).name() + + ". Skipping rest of data."); + } + + // Length of attached data is prepended as plaintext + final byte[] lengthBytes = new byte[4]; + Streams.readFully(mInput, lengthBytes); + final int length = ByteBuffer.wrap(lengthBytes).getInt(); + + // Read data based on the length + final byte[] data; + try { + data = new byte[length]; + } catch (OutOfMemoryError error) { + throw new SecureChannelException("Payload is too large.", error); + } + + Streams.readFully(mInput, data); + if (!MessageType.shouldEncrypt(expected)) { + return data; + } + + return mConnectionContext.decodeMessageFromPeer(data, headerBytes); } - - if (type != expected.mValue) { - Streams.skipByReading(mInput, Long.MAX_VALUE); - throw new SecureChannelException("Unexpected message type. Expected " + expected.name() - + "; Found " + MessageType.from(type).name() + ". Skipping rest of data."); - } - - // Length of attached data is prepended as plaintext - final byte[] lengthBytes = new byte[4]; - Streams.readFully(mInput, lengthBytes); - final int length = ByteBuffer.wrap(lengthBytes).getInt(); - - // Read data based on the length - final byte[] data; - try { - data = new byte[length]; - } catch (OutOfMemoryError error) { - throw new SecureChannelException("Payload is too large.", error); - } - - Streams.readFully(mInput, data); - if (!MessageType.shouldEncrypt(expected)) { - return data; - } - - return mConnectionContext.decodeMessageFromPeer(data, headerBytes); } - private void sendMessage(MessageType messageType, byte[] payload) + private void sendMessage(MessageType messageType, final byte[] payload) throws IOException, BadHandleException { synchronized (mOutput) { - byte[] header = ByteBuffer.allocate(HEADER_LENGTH) + final byte[] header = ByteBuffer.allocate(HEADER_LENGTH) .putInt(VERSION) .putShort(messageType.mValue) .array(); - byte[] data = MessageType.shouldEncrypt(messageType) + final byte[] data = MessageType.shouldEncrypt(messageType) ? mConnectionContext.encodeMessageToPeer(payload, header) : payload; mOutput.write(header); diff --git a/services/companion/java/com/android/server/companion/transport/CompanionTransportManager.java b/services/companion/java/com/android/server/companion/transport/CompanionTransportManager.java index 539020519f84a..092eb4ea9014e 100644 --- a/services/companion/java/com/android/server/companion/transport/CompanionTransportManager.java +++ b/services/companion/java/com/android/server/companion/transport/CompanionTransportManager.java @@ -46,6 +46,7 @@ import com.android.server.companion.AssociationStore; import java.io.IOException; import java.nio.ByteBuffer; +import java.nio.charset.StandardCharsets; import java.util.ArrayList; import java.util.List; import java.util.concurrent.CompletableFuture; @@ -296,26 +297,32 @@ public class CompanionTransportManager { Slog.i(TAG, "Remote device SDK: " + remoteSdk + ", release:" + new String(remoteRelease)); Transport transport = mTempTransport; - mTempTransport = null; + mTempTransport.stop(); int sdk = Build.VERSION.SDK_INT; String release = Build.VERSION.RELEASE; - if (remoteSdk == NON_ANDROID) { + if (Build.isDebuggable()) { + // Debug builds cannot pass attestation verification. Use hardcoded key instead. + Slog.d(TAG, "Creating an unauthenticated secure channel"); + final byte[] testKey = "CDM".getBytes(StandardCharsets.UTF_8); + transport = new SecureTransport(transport.getAssociationId(), transport.getFd(), + mContext, testKey, null); + } else if (remoteSdk == NON_ANDROID) { // TODO: pass in a real preSharedKey transport = new SecureTransport(transport.getAssociationId(), transport.getFd(), - mContext, null, null); - } else if (sdk < SECURE_CHANNEL_AVAILABLE_SDK - || remoteSdk < SECURE_CHANNEL_AVAILABLE_SDK) { - // TODO: depending on the release version, either - // 1) using a RawTransport for old T versions - // 2) or an Ukey2 handshaked transport for UKey2 backported T versions - } else { + mContext, new byte[0], null); + } else if (sdk >= SECURE_CHANNEL_AVAILABLE_SDK + && remoteSdk >= SECURE_CHANNEL_AVAILABLE_SDK) { Slog.i(TAG, "Creating a secure channel"); transport = new SecureTransport(transport.getAssociationId(), transport.getFd(), mContext); - addMessageListenersToTransport(transport); - transport.start(); + } else { + // TODO: depending on the release version, either + // 1) using a RawTransport for old T versions + // 2) or an Ukey2 handshaked transport for UKey2 backported T versions } + addMessageListenersToTransport(transport); + transport.start(); mTransports.put(transport.getAssociationId(), transport); // Doesn't need to notifyTransportsChanged here, it'll be done in attachSystemDataTransport } diff --git a/services/companion/java/com/android/server/companion/transport/RawTransport.java b/services/companion/java/com/android/server/companion/transport/RawTransport.java index 4060f6efe0cad..41589018b1494 100644 --- a/services/companion/java/com/android/server/companion/transport/RawTransport.java +++ b/services/companion/java/com/android/server/companion/transport/RawTransport.java @@ -36,6 +36,9 @@ class RawTransport extends Transport { @Override public void start() { + if (DEBUG) { + Slog.d(TAG, "Starting raw transport."); + } new Thread(() -> { try { while (!mStopped) { @@ -44,7 +47,7 @@ class RawTransport extends Transport { } catch (IOException e) { if (!mStopped) { Slog.w(TAG, "Trouble during transport", e); - stop(); + close(); } } }).start(); @@ -52,8 +55,19 @@ class RawTransport extends Transport { @Override public void stop() { + if (DEBUG) { + Slog.d(TAG, "Stopping raw transport."); + } mStopped = true; + } + @Override + public void close() { + stop(); + + if (DEBUG) { + Slog.d(TAG, "Closing raw transport."); + } IoUtils.closeQuietly(mRemoteIn); IoUtils.closeQuietly(mRemoteOut); } @@ -79,15 +93,17 @@ class RawTransport extends Transport { } private void receiveMessage() throws IOException { - final byte[] headerBytes = new byte[HEADER_LENGTH]; - Streams.readFully(mRemoteIn, headerBytes); - final ByteBuffer header = ByteBuffer.wrap(headerBytes); - final int message = header.getInt(); - final int sequence = header.getInt(); - final int length = header.getInt(); - final byte[] data = new byte[length]; - Streams.readFully(mRemoteIn, data); + synchronized (mRemoteIn) { + final byte[] headerBytes = new byte[HEADER_LENGTH]; + Streams.readFully(mRemoteIn, headerBytes); + final ByteBuffer header = ByteBuffer.wrap(headerBytes); + final int message = header.getInt(); + final int sequence = header.getInt(); + final int length = header.getInt(); + final byte[] data = new byte[length]; + Streams.readFully(mRemoteIn, data); - handleMessage(message, sequence, data); + handleMessage(message, sequence, data); + } } } diff --git a/services/companion/java/com/android/server/companion/transport/SecureTransport.java b/services/companion/java/com/android/server/companion/transport/SecureTransport.java index cca08435c0a5b..4054fc95f04ae 100644 --- a/services/companion/java/com/android/server/companion/transport/SecureTransport.java +++ b/services/companion/java/com/android/server/companion/transport/SecureTransport.java @@ -21,6 +21,7 @@ import android.content.Context; import android.os.ParcelFileDescriptor; import android.util.Slog; +import com.android.internal.annotations.GuardedBy; import com.android.server.companion.securechannel.AttestationVerifier; import com.android.server.companion.securechannel.SecureChannel; @@ -35,6 +36,7 @@ class SecureTransport extends Transport implements SecureChannel.Callback { private volatile boolean mShouldProcessRequests = false; + @GuardedBy("mRequestQueue") private final BlockingQueue mRequestQueue = new ArrayBlockingQueue<>(100); SecureTransport(int associationId, ParcelFileDescriptor fd, Context context) { @@ -59,6 +61,12 @@ class SecureTransport extends Transport implements SecureChannel.Callback { mShouldProcessRequests = false; } + @Override + public void close() { + mSecureChannel.close(); + mShouldProcessRequests = false; + } + @Override public Future requestForResponse(int message, byte[] data) { // Check if channel is secured and start securing @@ -85,12 +93,14 @@ class SecureTransport extends Transport implements SecureChannel.Callback { } // Queue up a message to send - mRequestQueue.add(ByteBuffer.allocate(HEADER_LENGTH + data.length) - .putInt(message) - .putInt(sequence) - .putInt(data.length) - .put(data) - .array()); + synchronized (mRequestQueue) { + mRequestQueue.add(ByteBuffer.allocate(HEADER_LENGTH + data.length) + .putInt(message) + .putInt(sequence) + .putInt(data.length) + .put(data) + .array()); + } } @Override @@ -102,9 +112,11 @@ class SecureTransport extends Transport implements SecureChannel.Callback { new Thread(() -> { try { while (mShouldProcessRequests) { - byte[] request = mRequestQueue.poll(); - if (request != null) { - mSecureChannel.sendSecureMessage(request); + synchronized (mRequestQueue) { + byte[] request = mRequestQueue.poll(); + if (request != null) { + mSecureChannel.sendSecureMessage(request); + } } } } catch (IOException e) { diff --git a/services/companion/java/com/android/server/companion/transport/Transport.java b/services/companion/java/com/android/server/companion/transport/Transport.java index d69ce8909c74c..d30104a095cfa 100644 --- a/services/companion/java/com/android/server/companion/transport/Transport.java +++ b/services/companion/java/com/android/server/companion/transport/Transport.java @@ -110,13 +110,26 @@ public abstract class Transport { return mFd; } + /** + * Start listening to messages. + */ public abstract void start(); + + /** + * Soft stop listening to the incoming data without closing the streams. + */ public abstract void stop(); + + /** + * Stop listening to the incoming data and close the streams. + */ + public abstract void close(); + protected abstract void sendMessage(int message, int sequence, @NonNull byte[] data) throws IOException; /** - * Send a message + * Send a message. */ public void sendMessage(int message, @NonNull byte[] data) throws IOException { sendMessage(message, mNextSequence.incrementAndGet(), data); @@ -170,7 +183,11 @@ public abstract class Transport { sendMessage(MESSAGE_RESPONSE_SUCCESS, sequence, data); break; } - case MESSAGE_REQUEST_PLATFORM_INFO: + case MESSAGE_REQUEST_PLATFORM_INFO: { + callback(message, data); + // DO NOT SEND A RESPONSE! + break; + } case MESSAGE_REQUEST_CONTEXT_SYNC: { callback(message, data); sendMessage(MESSAGE_RESPONSE_SUCCESS, sequence, EmptyArray.BYTE);