* PROJECT: ReactOS Kernel
* LICENSE: GPL - See COPYING in the top level directory
* FILE: ntoskrnl/lpc/connect.c
* PURPOSE: Local Procedure Call: Connection Management
* PROGRAMMERS: Alex Ionescu (alex.ionescu@reactos.org)
*/
#include <ntoskrnl.h>
#define NDEBUG
#include <debug.h>
PVOID
NTAPI
LpcpFreeConMsg(IN OUT PLPCP_MESSAGE *Message,
IN OUT PLPCP_CONNECTION_MESSAGE *ConnectMessage,
IN PETHREAD CurrentThread)
{
PVOID SectionToMap;
PLPCP_MESSAGE ReplyMessage;
KeAcquireGuardedMutex(&LpcpLock);
if (!IsListEmpty(&CurrentThread->LpcReplyChain))
{
RemoveEntryList(&CurrentThread->LpcReplyChain);
InitializeListHead(&CurrentThread->LpcReplyChain);
}
ReplyMessage = LpcpGetMessageFromThread(CurrentThread);
if (ReplyMessage)
{
*Message = ReplyMessage;
if (!IsListEmpty(&ReplyMessage->Entry))
{
RemoveEntryList(&ReplyMessage->Entry);
InitializeListHead(&ReplyMessage->Entry);
}
CurrentThread->LpcReceivedMessageId = 0;
CurrentThread->LpcReplyMessage = NULL;
*ConnectMessage = (PLPCP_CONNECTION_MESSAGE)(ReplyMessage + 1);
SectionToMap = (*ConnectMessage)->SectionToMap;
(*ConnectMessage)->SectionToMap = NULL;
}
else
{
*Message = NULL;
SectionToMap = NULL;
}
KeReleaseGuardedMutex(&LpcpLock);
return SectionToMap;
}
* @implemented
*/
NTSTATUS
NTAPI
NtSecureConnectPort(OUT PHANDLE PortHandle,
IN PUNICODE_STRING PortName,
IN PSECURITY_QUALITY_OF_SERVICE SecurityQos,
IN OUT PPORT_VIEW ClientView OPTIONAL,
IN PSID ServerSid OPTIONAL,
IN OUT PREMOTE_PORT_VIEW ServerView OPTIONAL,
OUT PULONG MaxMessageLength OPTIONAL,
IN OUT PVOID ConnectionInformation OPTIONAL,
IN OUT PULONG ConnectionInformationLength OPTIONAL)
{
NTSTATUS Status = STATUS_SUCCESS;
KPROCESSOR_MODE PreviousMode = KeGetPreviousMode();
PETHREAD Thread = PsGetCurrentThread();
#if DBG
UNICODE_STRING CapturedPortName;
#endif
SECURITY_QUALITY_OF_SERVICE CapturedQos;
PORT_VIEW CapturedClientView;
PSID CapturedServerSid;
ULONG ConnectionInfoLength = 0;
PLPCP_PORT_OBJECT Port, ClientPort;
PLPCP_MESSAGE Message;
PLPCP_CONNECTION_MESSAGE ConnectMessage;
ULONG PortMessageLength;
HANDLE Handle;
PVOID SectionToMap;
LARGE_INTEGER SectionOffset;
PTOKEN Token;
PTOKEN_USER TokenUserInfo;
PAGED_CODE();
if (PreviousMode != KernelMode)
{
_SEH2_TRY
{
ProbeForWriteHandle(PortHandle);
ProbeForRead(SecurityQos, sizeof(*SecurityQos), sizeof(ULONG));
CapturedQos = *(volatile SECURITY_QUALITY_OF_SERVICE*)SecurityQos;
if (ClientView)
{
ProbeForWrite(ClientView, sizeof(*ClientView), sizeof(ULONG));
CapturedClientView = *(volatile PORT_VIEW*)ClientView;
if (CapturedClientView.Length != sizeof(CapturedClientView))
{
_SEH2_YIELD(return STATUS_INVALID_PARAMETER);
}
}
if (ServerView)
{
ProbeForWrite(ServerView, sizeof(*ServerView), sizeof(ULONG));
if (((volatile REMOTE_PORT_VIEW*)ServerView)->Length != sizeof(*ServerView))
{
_SEH2_YIELD(return STATUS_INVALID_PARAMETER);
}
}
if (MaxMessageLength)
ProbeForWriteUlong(MaxMessageLength);
if (ConnectionInformationLength)
{
ProbeForWriteUlong(ConnectionInformationLength);
ConnectionInfoLength = *(volatile ULONG*)ConnectionInformationLength;
}
if (ConnectionInformation)
ProbeForWrite(ConnectionInformation, ConnectionInfoLength, sizeof(ULONG));
CapturedServerSid = ServerSid;
if (ServerSid)
{
Status = SepCaptureSid(ServerSid,
PreviousMode,
PagedPool,
TRUE,
&CapturedServerSid);
if (!NT_SUCCESS(Status))
{
DPRINT1("Failed to capture ServerSid!\n");
_SEH2_YIELD(return Status);
}
}
}
_SEH2_EXCEPT(EXCEPTION_EXECUTE_HANDLER)
{
_SEH2_YIELD(return _SEH2_GetExceptionCode());
}
_SEH2_END;
}
else
{
CapturedQos = *SecurityQos;
if (ClientView)
{
if (ClientView->Length != sizeof(*ClientView))
{
return STATUS_INVALID_PARAMETER;
}
CapturedClientView = *ClientView;
}
if (ServerView)
{
if (ServerView->Length != sizeof(*ServerView))
{
return STATUS_INVALID_PARAMETER;
}
}
if (ConnectionInformationLength)
ConnectionInfoLength = *ConnectionInformationLength;
CapturedServerSid = ServerSid;
}
#if DBG
* its own capture. As it is used only for debugging, ignore any failure;
* the string is zeroed out in such case. */
ProbeAndCaptureUnicodeString(&CapturedPortName, PreviousMode, PortName);
LPCTRACE(LPC_CONNECT_DEBUG,
"Name: %wZ. SecurityQos: %p. Views: %p/%p. Sid: %p\n",
&CapturedPortName,
SecurityQos,
ClientView,
ServerView,
ServerSid);
#endif
Status = ObReferenceObjectByName(PortName,
0,
NULL,
PORT_CONNECT,
LpcPortObjectType,
PreviousMode,
NULL,
(PVOID*)&Port);
if (!NT_SUCCESS(Status))
{
#if DBG
DPRINT1("Failed to reference port '%wZ': 0x%lx\n", &CapturedPortName, Status);
ReleaseCapturedUnicodeString(&CapturedPortName, PreviousMode);
#endif
if (CapturedServerSid != ServerSid)
SepReleaseSid(CapturedServerSid, PreviousMode, TRUE);
return Status;
}
if ((Port->Flags & LPCP_PORT_TYPE_MASK) != LPCP_CONNECTION_PORT)
{
#if DBG
DPRINT1("Port '%wZ' is not a connection port (Flags: 0x%lx)\n", &CapturedPortName, Port->Flags);
ReleaseCapturedUnicodeString(&CapturedPortName, PreviousMode);
#endif
ObDereferenceObject(Port);
if (CapturedServerSid != ServerSid)
SepReleaseSid(CapturedServerSid, PreviousMode, TRUE);
return STATUS_INVALID_PORT_HANDLE;
}
if (ServerSid)
{
if (Port->ServerProcess)
{
Token = PsReferencePrimaryToken(Port->ServerProcess);
Status = SeQueryInformationToken(Token, TokenUser, (PVOID*)&TokenUserInfo);
PsDereferencePrimaryToken(Token);
if (NT_SUCCESS(Status))
{
if (!RtlEqualSid(CapturedServerSid, TokenUserInfo->User.Sid))
{
#if DBG
DPRINT1("Port '%wZ': server SID mismatch\n", &CapturedPortName);
#endif
Status = STATUS_SERVER_SID_MISMATCH;
}
ExFreePoolWithTag(TokenUserInfo, TAG_SE);
}
}
else
{
#if DBG
DPRINT1("Port '%wZ': server SID mismatch\n", &CapturedPortName);
#endif
Status = STATUS_SERVER_SID_MISMATCH;
}
if (CapturedServerSid != ServerSid)
SepReleaseSid(CapturedServerSid, PreviousMode, TRUE);
}
#if DBG
ReleaseCapturedUnicodeString(&CapturedPortName, PreviousMode);
#endif
if (ServerSid && !NT_SUCCESS(Status))
{
ObDereferenceObject(Port);
return Status;
}
Status = ObCreateObject(PreviousMode,
LpcPortObjectType,
NULL,
PreviousMode,
NULL,
sizeof(LPCP_PORT_OBJECT),
0,
0,
(PVOID*)&ClientPort);
if (!NT_SUCCESS(Status))
{
DPRINT1("Failed to create Port object: 0x%lx\n", Status);
ObDereferenceObject(Port);
return Status;
}
* Setup the client port -- From now on, dereferencing the client port
* will automatically dereference the connection port too.
*/
RtlZeroMemory(ClientPort, sizeof(LPCP_PORT_OBJECT));
ClientPort->Flags = LPCP_CLIENT_PORT;
ClientPort->ConnectionPort = Port;
ClientPort->MaxMessageLength = Port->MaxMessageLength;
ClientPort->SecurityQos = CapturedQos;
InitializeListHead(&ClientPort->LpcReplyChainHead);
InitializeListHead(&ClientPort->LpcDataInfoChainHead);
if (CapturedQos.ContextTrackingMode == SECURITY_DYNAMIC_TRACKING)
{
ClientPort->Flags |= LPCP_SECURITY_DYNAMIC;
}
else
{
Status = SeCreateClientSecurity(Thread,
&CapturedQos,
FALSE,
&ClientPort->StaticSecurity);
if (!NT_SUCCESS(Status))
{
DPRINT1("SeCreateClientSecurity failed: 0x%lx\n", Status);
ObDereferenceObject(ClientPort);
return Status;
}
}
Status = LpcpInitializePortQueue(ClientPort);
if (!NT_SUCCESS(Status))
{
DPRINT1("LpcpInitializePortQueue failed: 0x%lx\n", Status);
ObDereferenceObject(ClientPort);
return Status;
}
if (ClientView)
{
Status = ObReferenceObjectByHandle(CapturedClientView.SectionHandle,
SECTION_MAP_READ |
SECTION_MAP_WRITE,
MmSectionObjectType,
PreviousMode,
(PVOID*)&SectionToMap,
NULL);
if (!NT_SUCCESS(Status))
{
DPRINT1("Failed to reference port section handle: 0x%lx\n", Status);
ObDereferenceObject(ClientPort);
return Status;
}
SectionOffset.QuadPart = CapturedClientView.SectionOffset;
Status = MmMapViewOfSection(SectionToMap,
PsGetCurrentProcess(),
&ClientPort->ClientSectionBase,
0,
0,
&SectionOffset,
&CapturedClientView.ViewSize,
ViewUnmap,
0,
PAGE_READWRITE);
CapturedClientView.SectionOffset = SectionOffset.LowPart;
if (!NT_SUCCESS(Status))
{
DPRINT1("Failed to map port section: 0x%lx\n", Status);
ObDereferenceObject(SectionToMap);
ObDereferenceObject(ClientPort);
return Status;
}
CapturedClientView.ViewBase = ClientPort->ClientSectionBase;
ClientPort->MappingProcess = PsGetCurrentProcess();
ObReferenceObject(ClientPort->MappingProcess);
}
else
{
SectionToMap = NULL;
}
if (ConnectionInfoLength > Port->MaxConnectionInfoLength)
{
ConnectionInfoLength = Port->MaxConnectionInfoLength;
}
Message = LpcpAllocateFromPortZone();
if (!Message)
{
DPRINT1("LpcpAllocateFromPortZone failed\n");
if (SectionToMap) ObDereferenceObject(SectionToMap);
ObDereferenceObject(ClientPort);
return STATUS_NO_MEMORY;
}
ConnectMessage = (PLPCP_CONNECTION_MESSAGE)(Message + 1);
Message->Request.ClientId = Thread->Cid;
if (ClientView)
{
Message->Request.ClientViewSize = CapturedClientView.ViewSize;
RtlCopyMemory(&ConnectMessage->ClientView,
&CapturedClientView,
sizeof(CapturedClientView));
RtlZeroMemory(&ConnectMessage->ServerView, sizeof(REMOTE_PORT_VIEW));
}
else
{
Message->Request.ClientViewSize = 0;
RtlZeroMemory(ConnectMessage, sizeof(LPCP_CONNECTION_MESSAGE));
}
ConnectMessage->ClientPort = NULL;
ConnectMessage->SectionToMap = SectionToMap;
Message->Request.u1.s1.DataLength = (CSHORT)ConnectionInfoLength +
sizeof(LPCP_CONNECTION_MESSAGE);
Message->Request.u1.s1.TotalLength = sizeof(LPCP_MESSAGE) +
Message->Request.u1.s1.DataLength;
Message->Request.u2.s2.Type = LPC_CONNECTION_REQUEST;
if (ConnectionInformation)
{
_SEH2_TRY
{
RtlCopyMemory(ConnectMessage + 1,
ConnectionInformation,
ConnectionInfoLength);
}
_SEH2_EXCEPT(EXCEPTION_EXECUTE_HANDLER)
{
DPRINT1("Exception 0x%lx when copying connection info to user mode\n",
_SEH2_GetExceptionCode());
LpcpFreeToPortZone(Message, 0);
if (SectionToMap) ObDereferenceObject(SectionToMap);
ObDereferenceObject(ClientPort);
_SEH2_YIELD(return _SEH2_GetExceptionCode());
}
_SEH2_END;
}
Status = STATUS_SUCCESS;
KeAcquireGuardedMutex(&LpcpLock);
if (Port->Flags & LPCP_NAME_DELETED)
{
Status = STATUS_OBJECT_NAME_NOT_FOUND;
}
else
{
Message->RepliedToThread = NULL;
Message->Request.MessageId = LpcpNextMessageId++;
if (!LpcpNextMessageId) LpcpNextMessageId = 1;
Thread->LpcReplyMessageId = Message->Request.MessageId;
InsertTailList(&Port->MsgQueue.ReceiveHead, &Message->Entry);
InsertTailList(&Port->LpcReplyChainHead, &Thread->LpcReplyChain);
Thread->LpcReplyMessage = Message;
ObReferenceObject(ClientPort);
ConnectMessage->ClientPort = ClientPort;
KeEnterCriticalRegion();
}
ObReferenceObject(Port);
KeReleaseGuardedMutex(&LpcpLock);
if (NT_SUCCESS(Status))
{
LPCTRACE(LPC_CONNECT_DEBUG,
"Messages: %p/%p. Ports: %p/%p. Status: %lx\n",
Message,
ConnectMessage,
Port,
ClientPort,
Status);
if (Port->Flags & LPCP_WAITABLE_PORT)
KeSetEvent(&Port->WaitEvent, 1, FALSE);
LpcpCompleteWait(Port->MsgQueue.Semaphore);
KeLeaveCriticalRegion();
LpcpConnectWait(&Thread->LpcReplySemaphore, PreviousMode);
}
SectionToMap = LpcpFreeConMsg(&Message, &ConnectMessage, Thread);
if (!NT_SUCCESS(Status))
{
if (KeReadStateSemaphore(&Thread->LpcReplySemaphore))
{
KeWaitForSingleObject(&Thread->LpcReplySemaphore,
WrExecutive,
KernelMode,
FALSE,
NULL);
}
goto Failure;
}
if (Message)
{
if ((Message->Request.u1.s1.DataLength -
sizeof(LPCP_CONNECTION_MESSAGE)) < ConnectionInfoLength)
{
ConnectionInfoLength = Message->Request.u1.s1.DataLength -
sizeof(LPCP_CONNECTION_MESSAGE);
}
if (ConnectionInformation)
{
_SEH2_TRY
{
if (ConnectionInformationLength)
*ConnectionInformationLength = ConnectionInfoLength;
RtlCopyMemory(ConnectionInformation,
ConnectMessage + 1,
ConnectionInfoLength);
}
_SEH2_EXCEPT(EXCEPTION_EXECUTE_HANDLER)
{
Status = _SEH2_GetExceptionCode();
_SEH2_YIELD(goto Failure);
}
_SEH2_END;
}
if (ClientPort->ConnectedPort)
{
PortMessageLength = Port->MaxMessageLength;
Status = ObInsertObject(ClientPort,
NULL,
PORT_ALL_ACCESS,
0,
NULL,
&Handle);
if (NT_SUCCESS(Status))
{
LPCTRACE(LPC_CONNECT_DEBUG,
"Handle: %p. Length: %lx\n",
Handle,
PortMessageLength);
_SEH2_TRY
{
*PortHandle = Handle;
if (MaxMessageLength)
*MaxMessageLength = PortMessageLength;
if (ClientView)
{
RtlCopyMemory(ClientView,
&ConnectMessage->ClientView,
sizeof(*ClientView));
}
if (ServerView)
{
RtlCopyMemory(ServerView,
&ConnectMessage->ServerView,
sizeof(*ServerView));
}
}
_SEH2_EXCEPT(EXCEPTION_EXECUTE_HANDLER)
{
ObCloseHandle(Handle, PreviousMode);
Status = _SEH2_GetExceptionCode();
}
_SEH2_END;
}
}
else
{
if (SectionToMap) ObDereferenceObject(SectionToMap);
KeAcquireGuardedMutex(&LpcpLock);
if (!(ClientPort->ConnectionPort) ||
(Port->Flags & LPCP_NAME_DELETED))
{
Status = STATUS_OBJECT_NAME_NOT_FOUND;
}
else
{
Status = STATUS_PORT_CONNECTION_REFUSED;
}
KeReleaseGuardedMutex(&LpcpLock);
ObDereferenceObject(ClientPort);
}
LpcpFreeToPortZone(Message, 0);
}
else
{
Status = STATUS_PORT_CONNECTION_REFUSED;
goto Failure;
}
ObDereferenceObject(Port);
return Status;
Failure:
if (Message) LpcpFreeToPortZone(Message, 0);
if (SectionToMap) ObDereferenceObject(SectionToMap);
ObDereferenceObject(ClientPort);
ObDereferenceObject(Port);
return Status;
}
* @implemented
*/
NTSTATUS
NTAPI
NtConnectPort(OUT PHANDLE PortHandle,
IN PUNICODE_STRING PortName,
IN PSECURITY_QUALITY_OF_SERVICE SecurityQos,
IN OUT PPORT_VIEW ClientView OPTIONAL,
IN OUT PREMOTE_PORT_VIEW ServerView OPTIONAL,
OUT PULONG MaxMessageLength OPTIONAL,
IN OUT PVOID ConnectionInformation OPTIONAL,
IN OUT PULONG ConnectionInformationLength OPTIONAL)
{
return NtSecureConnectPort(PortHandle,
PortName,
SecurityQos,
ClientView,
NULL,
ServerView,
MaxMessageLength,
ConnectionInformation,
ConnectionInformationLength);
}