diff --git a/CompiledServer.vcxproj b/CompiledServer.vcxproj
index 87483daf..499f4227 100644
--- a/CompiledServer.vcxproj
+++ b/CompiledServer.vcxproj
@@ -304,6 +304,7 @@
+
@@ -352,6 +353,7 @@
+
diff --git a/CompiledServer.vcxproj.filters b/CompiledServer.vcxproj.filters
index b6a59217..6285e716 100644
--- a/CompiledServer.vcxproj.filters
+++ b/CompiledServer.vcxproj.filters
@@ -159,6 +159,9 @@
Source Files
+
+ Source Files
+
@@ -374,5 +377,8 @@
Header Files
+
+ Header Files
+
\ No newline at end of file
diff --git a/Interface/Server.h b/Interface/Server.h
index c0998ea9..4a95c0bf 100644
--- a/Interface/Server.h
+++ b/Interface/Server.h
@@ -143,6 +143,7 @@ public:
virtual void StartCustomStreamService(IService *pService, std::string pServiceName, unsigned short pPort, int pMaxClientsPerThread=-1, BindTarget bindTarget=BindTarget_All)=0;
virtual IPipe* ConnectStream(std::string pServer, unsigned short pPort, unsigned int pTimeoutms=0)=0;
+ virtual IPipe* ConnectSslStream(const std::string& pServer, unsigned short pPort, unsigned int pTimeoutms) = 0;
virtual IPipe *PipeFromSocket(SOCKET pSocket)=0;
virtual void DisconnectStream(IPipe *pipe)=0;
virtual std::string LookupHostname(const std::string& pIp)=0;
diff --git a/SChannelPipe.cpp b/SChannelPipe.cpp
new file mode 100644
index 00000000..8066aee4
--- /dev/null
+++ b/SChannelPipe.cpp
@@ -0,0 +1,560 @@
+#include "SChannelPipe.h"
+#include "Server.h"
+#include "stringtools.h"
+#include
+#include
+#include
+
+
+PSecurityFunctionTableW SChannelPipe::sec = NULL;
+
+SChannelPipe::SChannelPipe(CStreamPipe * bpipe)
+ : bpipe(bpipe), has_cred_handle(false),
+ has_ctxt_handle(false), decbuf_pos(0),
+ sendbuf_pos(0), last_flush_time(0),
+ has_error(false)
+{
+}
+
+SChannelPipe::~SChannelPipe()
+{
+ if (has_cred_handle)
+ {
+ sec->FreeCredentialsHandle(&cred_handle);
+ }
+
+ if (has_ctxt_handle)
+ {
+ sec->DeleteSecurityContext(&ctxt_handle);
+ }
+
+ delete bpipe;
+}
+
+bool SChannelPipe::ssl_connect(const std::string& p_hostname, int timeoutms)
+{
+ hostname = p_hostname;
+ int64 starttime = Server->getTimeMS();
+
+ SCHANNEL_CRED cred_data = {};
+
+ cred_data.dwVersion = SCHANNEL_CRED_VERSION;
+ cred_data.dwFlags = SCH_CRED_AUTO_CRED_VALIDATION | SCH_CRED_REVOCATION_CHECK_CHAIN;
+ cred_data.grbitEnabledProtocols = SP_PROT_TLS1_0_CLIENT | SP_PROT_TLS1_1_CLIENT | SP_PROT_TLS1_2_CLIENT;
+
+ HRESULT res = sec->AcquireCredentialsHandleW(NULL, UNISP_NAME_W,
+ SECPKG_CRED_OUTBOUND, NULL, &cred_data, NULL, NULL, &cred_handle, &time_stamp);
+
+ if (res != SEC_E_OK)
+ {
+ Server->Log("AcquireCredentialsHandleW failed with result " + convert((int64)res), LL_ERROR);
+ return false;
+ }
+
+ has_cred_handle = true;
+
+ SecBuffer outbuf = {};
+ outbuf.BufferType = SECBUFFER_EMPTY;
+ SecBufferDesc outbuf_desc;
+ outbuf_desc.ulVersion = SECBUFFER_VERSION;
+ outbuf_desc.cBuffers = 1;
+ outbuf_desc.pBuffers = &outbuf;
+
+ std::wstring hostname_w = Server->ConvertToWchar(hostname);
+
+ unsigned long flags = ISC_REQ_SEQUENCE_DETECT | ISC_REQ_REPLAY_DETECT |
+ ISC_REQ_CONFIDENTIALITY | ISC_REQ_ALLOCATE_MEMORY |
+ ISC_REQ_STREAM;
+
+ unsigned long ret_flags = 0;
+ res = sec->InitializeSecurityContextW(&cred_handle, NULL, const_cast(hostname_w.c_str()), flags,
+ 0, 0, NULL, 0, &ctxt_handle, &outbuf_desc, &ret_flags, &time_stamp);
+
+ if (res != SEC_I_CONTINUE_NEEDED)
+ {
+ Server->Log("InitializeSecurityContextW failed with result " + convert((int64)res), LL_ERROR);
+ return false;
+ }
+
+ has_ctxt_handle = true;
+
+ if (flags != ret_flags)
+ {
+ Server->Log("Setting security context flags failed " + convert((int64)flags) + "!=" + convert((int64)ret_flags), LL_ERROR);
+ return false;
+ }
+
+ if (!bpipe->Write(reinterpret_cast(outbuf.pvBuffer), outbuf.cbBuffer, timeoutms, true))
+ {
+ return false;
+ }
+
+ last_flush_time = Server->getTimeMS();
+ int64 passed_time = last_flush_time - starttime;
+ int64 remaining_time = timeoutms == -1 ? -1 : (passed_time < timeoutms ? (timeoutms - passed_time) : 0);
+ return ssl_connect_negotiate(static_cast(remaining_time), true);
+}
+
+bool SChannelPipe::ssl_connect_negotiate(int timeoutms, bool do_read)
+{
+ const size_t encbuf_size_incr = 4096;
+
+ bool connected = false;
+
+ int64 starttime = Server->getTimeMS();
+
+ unsigned long flags = ISC_REQ_SEQUENCE_DETECT | ISC_REQ_REPLAY_DETECT |
+ ISC_REQ_CONFIDENTIALITY | ISC_REQ_ALLOCATE_MEMORY |
+ ISC_REQ_STREAM;
+
+ unsigned long ret_flags;
+
+ std::wstring hostname_w = Server->ConvertToWchar(hostname);
+
+ int64 passed_time;
+ encbuf_pos = 0;
+ while (timeoutms == -1
+ || (passed_time = Server->getTimeMS() - starttime)getTimeMS() - starttime;
+
+ if (do_read)
+ {
+ do_read = false;
+
+ if (!bpipe->isReadable(timeoutms == -1 ? timeoutms : (timeoutms - passed_time)))
+ return false;
+
+ if (encbuf.size() - encbuf_pos < encbuf_size_incr)
+ {
+ encbuf.resize(encbuf.size() + encbuf_size_incr);
+ }
+ size_t read = bpipe->Read(&encbuf[encbuf_pos], encbuf_size_incr, 0);
+
+ if (read == 0)
+ return false;
+
+ encbuf_pos += read;
+ }
+
+ SecBuffer inbuf2[2] = {};
+ inbuf2[0].BufferType = SECBUFFER_TOKEN;
+ inbuf2[0].cbBuffer = encbuf_pos;
+ inbuf2[0].pvBuffer = encbuf.data();
+ inbuf2[1].BufferType = SECBUFFER_EMPTY;
+
+ SecBufferDesc inbuf2_desc;
+ inbuf2_desc.cBuffers = 2;
+ inbuf2_desc.pBuffers = inbuf2;
+ inbuf2_desc.ulVersion = SECBUFFER_VERSION;
+
+ SecBuffer outbuf[3] = {};
+ outbuf[0].BufferType = SECBUFFER_TOKEN;
+ outbuf[1].BufferType = SECBUFFER_ALERT;
+ outbuf[2].BufferType = SECBUFFER_EMPTY;
+
+ SecBufferDesc outbuf_desc;
+ outbuf_desc.cBuffers = 3;
+ outbuf_desc.pBuffers = outbuf;
+ outbuf_desc.ulVersion = SECBUFFER_VERSION;
+
+ HRESULT res = sec->InitializeSecurityContextW(&cred_handle, &ctxt_handle,
+ const_cast(hostname_w.c_str()), flags, 0, 0, &inbuf2_desc,
+ 0, NULL, &outbuf_desc, &ret_flags, &time_stamp);
+
+ if (res == SEC_E_INCOMPLETE_MESSAGE)
+ {
+ do_read = true;
+ continue;
+ }
+ else
+ {
+ if (inbuf2[1].BufferType == SECBUFFER_EXTRA
+ && inbuf2[1].cbBuffer>0)
+ {
+ encbuf.erase(encbuf.begin(), encbuf.begin() + (encbuf_pos - inbuf2[1].cbBuffer));
+ encbuf_pos -= encbuf_pos - inbuf2[1].cbBuffer;
+
+ if (res == SEC_I_CONTINUE_NEEDED)
+ {
+ continue;
+ }
+ }
+ else
+ {
+ encbuf.clear();
+ encbuf_pos = 0;
+ }
+ }
+
+ if (res == SEC_E_OK
+ || res == SEC_I_CONTINUE_NEEDED)
+ {
+ bool has_error = false;
+ for (size_t i = 0; i < outbuf_desc.cBuffers; ++i)
+ {
+ SecBuffer& buf = outbuf_desc.pBuffers[i];
+ if (buf.BufferType == SECBUFFER_TOKEN
+ && buf.cbBuffer > 0
+ && !has_error)
+ {
+ passed_time = Server->getTimeMS() - starttime;
+ int64 remaining_time = timeoutms == -1 ? -1 : (passed_time < timeoutms ? (timeoutms - passed_time) : 0);
+ if (!bpipe->Write(reinterpret_cast(buf.pvBuffer), buf.cbBuffer, static_cast(remaining_time), true))
+ {
+ has_error = true;
+ }
+ }
+
+ if (buf.pvBuffer != NULL)
+ sec->FreeContextBuffer(buf.pvBuffer);
+ }
+
+ if (has_error)
+ return false;
+
+ if (res == SEC_E_OK)
+ {
+ connected = true;
+ break;
+ }
+ else
+ {
+ do_read = true;
+ }
+ }
+ else
+ {
+ return false;
+ }
+ }
+
+ if (connected)
+ {
+ HRESULT res = sec->QueryContextAttributesW(&ctxt_handle, SECPKG_ATTR_STREAM_SIZES, &stream_sizes);
+
+ if (res != SEC_E_OK)
+ {
+ return false;
+ }
+
+ header_buf.resize(stream_sizes.cbHeader);
+ trailer_buf.resize(stream_sizes.cbTrailer);
+ }
+
+ return connected;
+}
+
+
+void SChannelPipe::init()
+{
+ HMODULE mod = LoadLibraryW(L"secur32.dll");
+
+ if (mod == NULL)
+ {
+ Server->Log("Error loading libarary secur32.dll. Errno: " + convert((int64)GetLastError()), LL_ERROR);
+ return;
+ }
+
+ INIT_SECURITY_INTERFACE_W init_sec = reinterpret_cast(GetProcAddress(mod, "InitSecurityInterfaceW"));
+
+ if (init_sec == NULL)
+ {
+ Server->Log("Error getting proc InitSecurityInterfaceW in secur32.dll. Errno: " + convert((int64)GetLastError()), LL_ERROR);
+ return;
+ }
+
+ sec = init_sec();
+}
+
+size_t SChannelPipe::Read(char * buffer, size_t bsize, int timeoutms)
+{
+ const size_t encbuf_size_incr = 1024;
+
+ if (has_error)
+ return 0;
+
+ if (decbuf_pos>0)
+ {
+ size_t toread = (std::min)(decbuf_pos, bsize);
+ memcpy(buffer, decbuf.data(), toread);
+ decbuf.erase(decbuf.begin(), decbuf.begin() + toread);
+ decbuf_pos -= toread;
+ return toread;
+ }
+
+ int64 starttime = Server->getTimeMS();
+
+ if (encbuf.size() - encbuf_pos < bsize)
+ {
+ encbuf.resize(encbuf.size() + bsize);
+ }
+
+ size_t read = bpipe->Read(&encbuf[encbuf_pos], bsize, timeoutms);
+
+ if (read == 0)
+ return 0;
+
+ encbuf_pos += read;
+
+ size_t orig_bsize = bsize;
+
+ HRESULT res = SEC_E_OK;
+ while ( (res == SEC_E_OK || res== SEC_E_INCOMPLETE_MESSAGE)
+ && encbuf_pos > 0 && bsize>0)
+ {
+ if (res == SEC_E_INCOMPLETE_MESSAGE)
+ {
+ if (encbuf.size() - encbuf_pos < encbuf_size_incr)
+ {
+ encbuf.resize(encbuf.size() + encbuf_size_incr);
+ }
+
+ int64 passed_time = Server->getTimeMS() - starttime;
+ int remaining_time = timeoutms == -1 ? -1 : ((timeoutms - passed_time) < 0 ? 0 : (timeoutms - passed_time));
+
+ size_t read = bpipe->Read(&encbuf[encbuf_pos], encbuf_size_incr, remaining_time);
+
+ if (read == 0)
+ return 0;
+
+ encbuf_pos += read;
+ }
+
+ SecBuffer inbuf[4] = {};
+ inbuf[0].BufferType = SECBUFFER_DATA;
+ inbuf[0].pvBuffer = encbuf.data();
+ inbuf[0].cbBuffer = encbuf_pos;
+ inbuf[1].BufferType = SECBUFFER_EMPTY;
+ inbuf[2].BufferType = SECBUFFER_EMPTY;
+ inbuf[3].BufferType = SECBUFFER_EMPTY;
+
+ SecBufferDesc inbuf_desc;
+ inbuf_desc.ulVersion = SECBUFFER_VERSION;
+ inbuf_desc.cBuffers = 4;
+ inbuf_desc.pBuffers = inbuf;
+
+ res = sec->DecryptMessage(&ctxt_handle, &inbuf_desc, 0, NULL);
+
+ if (res == SEC_E_OK
+ || res== SEC_I_RENEGOTIATE)
+ {
+ if (inbuf[1].BufferType == SECBUFFER_DATA
+ && inbuf[1].cbBuffer>0)
+ {
+ size_t toread = 0;
+ if (bsize > 0)
+ {
+ toread = (std::min)(bsize, (size_t)inbuf[1].cbBuffer);
+ memcpy(buffer, inbuf[1].pvBuffer, toread);
+ inbuf[1].cbBuffer -= toread;
+ bsize -= toread;
+ buffer += toread;
+ }
+
+ if (inbuf[1].cbBuffer > 0)
+ {
+ if (decbuf.size() - decbuf_pos < inbuf[1].cbBuffer)
+ {
+ decbuf.resize(decbuf.size() + inbuf[1].cbBuffer);
+ }
+
+ memcpy(&decbuf[decbuf_pos], reinterpret_cast(inbuf[1].pvBuffer) + toread, inbuf[1].cbBuffer);
+ decbuf_pos += inbuf[1].cbBuffer;
+ }
+ }
+
+ if (inbuf[3].BufferType == SECBUFFER_EXTRA
+ && inbuf[3].cbBuffer>0)
+ {
+ encbuf.erase(encbuf.begin(), encbuf.begin() + (encbuf_pos - inbuf[3].cbBuffer));
+ encbuf_pos -= encbuf_pos - inbuf[3].cbBuffer;
+ }
+ else
+ {
+ encbuf.clear();
+ encbuf_pos = 0;
+ }
+ }
+
+ if (res == SEC_I_RENEGOTIATE)
+ {
+ if (!ssl_connect_negotiate(timeoutms, false))
+ {
+ return 0;
+ }
+ }
+
+ if (res == SEC_I_CONTEXT_EXPIRED)
+ {
+ return 0;
+ }
+ }
+
+ return orig_bsize-bsize;
+}
+
+bool SChannelPipe::Write(const char * buffer, size_t bsize, int timeoutms, bool flush)
+{
+ if (has_error)
+ return false;
+
+ if (sendbuf_pos + bsize > sendbuf.size())
+ {
+ sendbuf.resize(sendbuf_pos + bsize);
+ }
+
+ memcpy(&sendbuf[sendbuf_pos], buffer, bsize);
+ sendbuf_pos += bsize;
+
+ if ((flush || sendbuf_pos>128 * 1024 || sendbuf_pos>= stream_sizes.cbMaximumMessage || (Server->getTimeMS() - last_flush_time)>200)
+ && sendbuf_pos>0)
+ {
+ return Flush(timeoutms);
+ }
+
+ return true;
+}
+
+size_t SChannelPipe::Read(std::string * ret, int timeoutms)
+{
+ if (has_error)
+ return 0;
+
+ char buf[1024];
+ size_t read = Read(buf, sizeof(buf), timeoutms);
+ if (read > 0)
+ ret->assign(buf, read);
+ return read;
+}
+
+bool SChannelPipe::Write(const std::string & str, int timeoutms, bool flush)
+{
+ return Write(str.data(), str.size(), timeoutms, flush);
+}
+
+bool SChannelPipe::Flush(int timeoutms)
+{
+ if (has_error)
+ return false;
+
+ size_t sendbuf_off = 0;
+ while (sendbuf_pos- sendbuf_off> 0)
+ {
+ size_t toflush = (std::min)((size_t)stream_sizes.cbMaximumMessage, sendbuf_pos- sendbuf_off);
+
+ SecBuffer outbuf[4] = {};
+ outbuf[0].BufferType = SECBUFFER_STREAM_HEADER;
+ outbuf[0].cbBuffer = header_buf.size();
+ outbuf[0].pvBuffer = header_buf.data();
+ outbuf[1].BufferType = SECBUFFER_DATA;
+ outbuf[1].cbBuffer = toflush;
+ outbuf[1].pvBuffer = &sendbuf[sendbuf_off];
+ outbuf[2].BufferType = SECBUFFER_STREAM_TRAILER;
+ outbuf[2].cbBuffer = trailer_buf.size();
+ outbuf[2].pvBuffer = trailer_buf.data();
+ outbuf[3].BufferType = SECBUFFER_EMPTY;
+
+ SecBufferDesc outbuf_desc;
+ outbuf_desc.ulVersion = SECBUFFER_VERSION;
+ outbuf_desc.cBuffers = 4;
+ outbuf_desc.pBuffers = outbuf;
+
+ HRESULT res = sec->EncryptMessage(&ctxt_handle, 0, &outbuf_desc, 0);
+
+ if (res != SEC_E_OK)
+ {
+ return false;
+ }
+
+ if (outbuf[0].cbBuffer > 0)
+ {
+ if (!bpipe->Write(reinterpret_cast(outbuf[0].pvBuffer), outbuf[0].cbBuffer, timeoutms, false))
+ {
+ has_error = true;
+ return false;
+ }
+ }
+
+ if (outbuf[1].cbBuffer > 0)
+ {
+ if (!bpipe->Write(reinterpret_cast(outbuf[1].pvBuffer), outbuf[1].cbBuffer, timeoutms, false))
+ {
+ has_error = true;
+ return false;
+ }
+ }
+
+ if (outbuf[2].cbBuffer > 0)
+ {
+ if (!bpipe->Write(reinterpret_cast(outbuf[2].pvBuffer), outbuf[2].cbBuffer, timeoutms, false))
+ {
+ has_error = true;
+ return false;
+ }
+ }
+
+ sendbuf_off += toflush;
+ }
+
+ sendbuf_pos -= sendbuf_off;
+
+ return bpipe->Flush(timeoutms);
+}
+
+bool SChannelPipe::isWritable(int timeoutms)
+{
+ if (has_error)
+ return false;
+
+ return bpipe->isWritable(timeoutms);
+}
+
+bool SChannelPipe::isReadable(int timeoutms)
+{
+ if (has_error)
+ return false;
+
+ return bpipe->isReadable(timeoutms);
+}
+
+bool SChannelPipe::hasError(void)
+{
+ return has_error || bpipe->hasError();
+}
+
+void SChannelPipe::shutdown(void)
+{
+ bpipe->shutdown();
+}
+
+size_t SChannelPipe::getNumElements(void)
+{
+ return bpipe->getNumElements();
+}
+
+void SChannelPipe::addThrottler(IPipeThrottler * throttler)
+{
+ bpipe->addThrottler(throttler);
+}
+
+void SChannelPipe::addOutgoingThrottler(IPipeThrottler * throttler)
+{
+ bpipe->addOutgoingThrottler(throttler);
+}
+
+void SChannelPipe::addIncomingThrottler(IPipeThrottler * throttler)
+{
+ bpipe->addIncomingThrottler(throttler);
+}
+
+_i64 SChannelPipe::getTransferedBytes(void)
+{
+ return bpipe->getTransferedBytes();
+}
+
+void SChannelPipe::resetTransferedBytes(void)
+{
+ bpipe->resetTransferedBytes();
+}
\ No newline at end of file
diff --git a/SChannelPipe.h b/SChannelPipe.h
new file mode 100644
index 00000000..d2bf5642
--- /dev/null
+++ b/SChannelPipe.h
@@ -0,0 +1,76 @@
+#pragma once
+
+#include "Interface/Pipe.h"
+#include "StreamPipe.h"
+#define SECURITY_WIN32
+#include
+#include
+
+class SChannelPipe : public IPipe
+{
+public:
+ SChannelPipe(CStreamPipe* bpipe);
+
+ ~SChannelPipe();
+
+ bool ssl_connect(const std::string& p_hostname, int timeoutms);
+
+ static void init();
+
+ // Inherited via IPipe
+ virtual size_t Read(char * buffer, size_t bsize, int timeoutms = -1);
+
+ virtual bool Write(const char * buffer, size_t bsize, int timeoutms = -1, bool flush = true);
+
+ virtual size_t Read(std::string * ret, int timeoutms = -1);
+
+ virtual bool Write(const std::string & str, int timeoutms = -1, bool flush = true);
+
+ virtual bool Flush(int timeoutms = -1);
+
+ virtual bool isWritable(int timeoutms = 0);
+
+ virtual bool isReadable(int timeoutms = 0);
+
+ virtual bool hasError(void);
+
+ virtual void shutdown(void);
+
+ virtual size_t getNumElements(void);
+
+ virtual void addThrottler(IPipeThrottler * throttler);
+
+ virtual void addOutgoingThrottler(IPipeThrottler * throttler);
+
+ virtual void addIncomingThrottler(IPipeThrottler * throttler);
+
+ virtual _i64 getTransferedBytes(void);
+
+ virtual void resetTransferedBytes(void);
+
+private:
+ bool ssl_connect_negotiate(int timeoutms, bool do_read);
+
+ CStreamPipe* bpipe;
+
+ static PSecurityFunctionTableW sec;
+
+ bool has_cred_handle;
+ CredHandle cred_handle;
+ TimeStamp time_stamp;
+ bool has_ctxt_handle;
+ CtxtHandle ctxt_handle;
+ std::vector encbuf;
+ size_t encbuf_pos;
+ std::vector decbuf;
+ size_t decbuf_pos;
+ std::vector sendbuf;
+ size_t sendbuf_pos;
+ std::string hostname;
+ int64 last_flush_time;
+ std::vector header_buf;
+ std::vector trailer_buf;
+ bool has_error;
+
+ SecPkgContext_StreamSizes stream_sizes;
+};
\ No newline at end of file
diff --git a/Server.cpp b/Server.cpp
index 2f3139cc..77111a25 100644
--- a/Server.cpp
+++ b/Server.cpp
@@ -57,8 +57,7 @@
#include "PipeThrottler.h"
#include "mt19937ar.h"
#include "Query.h"
-
-
+#include "SChannelPipe.h"
#ifdef _WIN32
#include
@@ -91,6 +90,7 @@
#include
#endif
+
const size_t SEND_BLOCKSIZE=8192;
const size_t MAX_THREAD_ID=std::string::npos;
@@ -207,6 +207,8 @@ void CServer::setup(void)
#ifdef MODE_WIN
File::init_mutex();
#endif
+
+ SChannelPipe::init();
}
void CServer::destroyAllDatabases(void)
@@ -1137,6 +1139,28 @@ IPipe* CServer::ConnectStream(std::string pServer, unsigned short pPort, unsigne
}
}
+IPipe * CServer::ConnectSslStream(const std::string & pServer, unsigned short pPort, unsigned int pTimeoutms)
+{
+ int64 starttime = Server->getTimeMS();
+ CStreamPipe* bpipe = static_cast(ConnectStream(pServer, pPort, pTimeoutms));
+
+ if (bpipe == NULL)
+ return NULL;
+
+ int64 remaining_time = pTimeoutms - (Server->getTimeMS() - starttime);
+ if (remaining_time < 0) remaining_time = 0;
+
+ SChannelPipe* ssl_pipe = new SChannelPipe(bpipe);
+
+ if (!ssl_pipe->ssl_connect(pServer, static_cast(remaining_time)))
+ {
+ delete ssl_pipe;
+ return NULL;
+ }
+
+ return ssl_pipe;
+}
+
IPipe *CServer::PipeFromSocket(SOCKET pSocket)
{
return new CStreamPipe(pSocket);
diff --git a/Server.h b/Server.h
index 1faac313..2b578bdd 100644
--- a/Server.h
+++ b/Server.h
@@ -137,6 +137,7 @@ public:
virtual void StartCustomStreamService(IService *pService, std::string pServiceName, unsigned short pPort, int pMaxClientsPerThread=-1, IServer::BindTarget bindTarget=IServer::BindTarget_All);
virtual IPipe* ConnectStream(std::string pServer, unsigned short pPort, unsigned int pTimeoutms);
+ virtual IPipe* ConnectSslStream(const std::string& pServer, unsigned short pPort, unsigned int pTimeoutms);
virtual IPipe *PipeFromSocket(SOCKET pSocket);
virtual void DisconnectStream(IPipe *pipe);
virtual std::string LookupHostname(const std::string& pIp);
diff --git a/StreamPipe.h b/StreamPipe.h
index 8a9a19ef..57a5b2b7 100644
--- a/StreamPipe.h
+++ b/StreamPipe.h
@@ -1,3 +1,5 @@
+#pragma once
+
#include "Interface/Pipe.h"
#include "socket_header.h"
#include
diff --git a/urbackupclient/InternetClient.cpp b/urbackupclient/InternetClient.cpp
index 80b0eb57..708cfc1b 100644
--- a/urbackupclient/InternetClient.cpp
+++ b/urbackupclient/InternetClient.cpp
@@ -38,6 +38,7 @@
#include
#include
+#include
#include "../cryptoplugin/ICryptoFactory.h"
@@ -221,7 +222,7 @@ void InternetClient::operator()(void)
{
if(n_connectionsgetThreadPool()->execute(new InternetClientThread(NULL, server_settings), "internet client");
+ Server->getThreadPool()->execute(new InternetClientThread(NULL, server_settings, NULL), "internet client");
newConnection();
}
else
@@ -278,6 +279,7 @@ void InternetClient::doUpdateSettings(void)
std::string server_name;
std::string computername;
std::string server_port="55415";
+ std::string server_proxy;
std::string authkey;
if(!settings->getValue("internet_authkey", &authkey) && !settings->getValue("internet_authkey_def", &authkey))
{
@@ -299,6 +301,8 @@ void InternetClient::doUpdateSettings(void)
{
computername=(IndexThread::getFileSrv()->getServerName());
}
+ if (!settings->getValue("internet_server_proxy", &server_proxy))
+ settings->getValue("internet_server_proxy_def", &server_proxy);
if( (settings->getValue("internet_server", &server_name) || settings->getValue("internet_server_def", &server_name))
&& !server_name.empty() )
{
@@ -311,18 +315,27 @@ void InternetClient::doUpdateSettings(void)
std::vector server_ports;
Tokenize(server_port, server_ports, ";");
+ std::vector server_proxies;
+ Tokenize(server_proxy, server_proxies, ";");
+
for(size_t i=0;i(atoi(server_ports[i].c_str()))));
+ connection_settings.port = static_cast(atoi(server_ports[i].c_str()));
}
- else
+ else if(!server_ports.empty())
{
- server_settings.servers.push_back(std::make_pair(server_names[i],
- static_cast(atoi(server_ports[server_ports.size()-1].c_str()))));
+ connection_settings.port = static_cast(atoi(server_ports[server_ports.size() - 1].c_str()));
}
+ if (i < server_proxies.size())
+ connection_settings.proxy = server_proxies[i];
+ else if(!server_proxies.empty())
+ connection_settings.proxy = server_proxies[server_proxies.size()-1];
+
+ server_settings.servers.push_back(connection_settings);
}
server_settings.clientname=computername;
server_settings.authkey=authkey;
@@ -340,6 +353,7 @@ void InternetClient::doUpdateSettings(void)
connected = false;
}
}
+
std::string tmp;
server_settings.internet_compress=true;
if(settings->getValue("internet_compress", &tmp) || settings->getValue("internet_compress_def", &tmp) )
@@ -378,20 +392,20 @@ bool InternetClient::tryToConnect(IScopedLock *lock)
for(size_t i=0;irelock(NULL);
- Server->Log("Trying to connect to internet server \""+name+"\" at port "+convert(port), LL_DEBUG);
- IPipe *cs=Server->ConnectStream(name, port, 10000);
+ Server->Log("Trying to connect to internet server \""+ selected_server_settings .hostname+"\" at port "+convert(selected_server_settings.port)
+ + (selected_server_settings.proxy.empty() ? "" : (" via HTTP proxy "+ selected_server_settings.proxy)), LL_DEBUG);
+ std::auto_ptr tcpstack(new CTCPStack);
+ IPipe *cs = connect(selected_server_settings, *tcpstack);
lock->relock(mutex);
if(cs!=NULL)
{
server_settings.selected_server=i;
Server->Log("Successfully connected.", LL_DEBUG);
setStatusMsg("connected");
- Server->getThreadPool()->execute(new InternetClientThread(cs, server_settings), "internet client");
+ Server->getThreadPool()->execute(new InternetClientThread(cs, server_settings, tcpstack.release()), "internet client");
newConnection();
return true;
}
@@ -486,9 +500,16 @@ void InternetClient::setStatusMsg(const std::string& msg)
status_msg=msg;
}
-InternetClientThread::InternetClientThread(IPipe *cs, const SServerSettings &server_settings)
- : cs(cs), server_settings(server_settings)
+InternetClientThread::InternetClientThread(IPipe *cs, const SServerSettings &server_settings, CTCPStack* tcpstack)
+ : cs(cs), server_settings(server_settings), tcpstack(tcpstack)
{
+ if (this->tcpstack == NULL)
+ this->tcpstack = new CTCPStack(true);
+}
+
+InternetClientThread::~InternetClientThread()
+{
+ delete tcpstack;
}
char *InternetClientThread::getReply(CTCPStack *tcpstack, IPipe *pipe, size_t &replysize, unsigned int timeoutms)
@@ -513,7 +534,6 @@ char *InternetClientThread::getReply(CTCPStack *tcpstack, IPipe *pipe, size_t &r
void InternetClientThread::operator()(void)
{
- CTCPStack tcpstack(true);
bool finish_ok=false;
bool rm_connection=true;
@@ -522,20 +542,20 @@ void InternetClientThread::operator()(void)
int tries=10;
while(tries>0 && cs==NULL)
{
- cs=Server->ConnectStream(server_settings.servers[server_settings.selected_server].first,
- server_settings.servers[server_settings.selected_server].second, 10000);
+ cs = InternetClient::connect(server_settings.servers[server_settings.selected_server], *tcpstack);
--tries;
InternetClient::setStatusMsg("connecting_failed");
if(cs==NULL && tries>0)
{
- Server->Log("Connecting to server "+server_settings.servers[server_settings.selected_server].first
- + " failed. Retrying in 30s...", LL_INFO);
+ Server->Log("Connecting to server "+server_settings.servers[server_settings.selected_server].hostname
+ +(server_settings.servers[server_settings.selected_server].proxy.empty() ? "" : (" via HTTP proxy " + server_settings.servers[server_settings.selected_server].proxy)) + " failed. Retrying in 30s...", LL_INFO);
Server->wait(30000);
}
}
if(cs==NULL)
{
- Server->Log("Connecting to server "+server_settings.servers[server_settings.selected_server].first
+ Server->Log("Connecting to server "+server_settings.servers[server_settings.selected_server].hostname
+ +(server_settings.servers[server_settings.selected_server].proxy.empty() ? "" : (" via HTTP proxy " + server_settings.servers[server_settings.selected_server].proxy))
+ " failed", LL_INFO);
InternetClient::rmConnection();
InternetClient::setHasConnection(false);
@@ -572,7 +592,7 @@ void InternetClientThread::operator()(void)
{
char *buf;
size_t bufsize;
- buf=getReply(&tcpstack, cs, bufsize, ic_auth_timeout);
+ buf=getReply(tcpstack, cs, bufsize, ic_auth_timeout);
if(buf==NULL)
{
Server->Log("Error receiving challenge packet");
@@ -663,7 +683,7 @@ void InternetClientThread::operator()(void)
data.addString(client_challenge);
data.addUInt(pbkdf2_iterations);
- tcpstack.Send(cs, data);
+ tcpstack->Send(cs, data);
challenge_response=crypto_fak->generateBinaryPasswordHash(hmac_key, client_challenge, 1);
}
@@ -671,7 +691,7 @@ void InternetClientThread::operator()(void)
{
char *buf;
size_t bufsize;
- buf=getReply(&tcpstack, cs, bufsize, ic_auth_timeout);
+ buf=getReply(tcpstack, cs, bufsize, ic_auth_timeout);
if(buf==NULL)
{
Server->Log("Error receiving authentication response");
@@ -746,7 +766,7 @@ void InternetClientThread::operator()(void)
data.addUInt(capa);
- tcpstack.Send(ics_pipe, data);
+ tcpstack->Send(ics_pipe, data);
}
comm_pipe=cs;
@@ -784,7 +804,7 @@ void InternetClientThread::operator()(void)
ping_timeout=ic_ping_timeout;
}
- buf=getReply(&tcpstack, comm_pipe, bufsize, ping_timeout);
+ buf=getReply(tcpstack, comm_pipe, bufsize, ping_timeout);
if(buf==NULL)
{
goto cleanup;
@@ -801,7 +821,7 @@ void InternetClientThread::operator()(void)
{
CWData data;
data.addChar(ID_ISC_PONG);
- tcpstack.Send(comm_pipe, data);
+ tcpstack->Send(comm_pipe, data);
}
else if(id==ID_ISC_CONNECT)
{
@@ -812,7 +832,7 @@ void InternetClientThread::operator()(void)
{
CWData data;
data.addChar(ID_ISC_CONNECT_OK);
- tcpstack.Send(comm_pipe, data);
+ tcpstack->Send(comm_pipe, data);
InternetClient::rmConnection();
rm_connection=false;
@@ -868,7 +888,7 @@ cleanup:
void InternetClientThread::runServiceWrapper(IPipe *pipe, ICustomClient *client)
{
- client->Init(Server->getThreadID(), pipe, server_settings.servers[server_settings.selected_server].first);
+ client->Init(Server->getThreadID(), pipe, server_settings.servers[server_settings.selected_server].hostname);
ClientConnector * cc=dynamic_cast(client);
if(cc!=NULL)
{
@@ -934,3 +954,167 @@ void InternetClientThread::printInfo( IPipe * pipe )
}
}
}
+
+IPipe * InternetClient::connect(const SServerConnectionSettings & selected_server_settings, CTCPStack& tcpstack)
+{
+ std::string proxy = selected_server_settings.proxy;
+ if (!proxy.empty())
+ {
+ bool ssl = false;
+ if (next(proxy, 0, "http://"))
+ {
+ ssl = false;
+ proxy = proxy.substr(7);
+ }
+ else if (next(proxy, 0, "https://"))
+ {
+ ssl = true;
+ proxy = proxy.substr(7);
+ }
+
+ std::string authorization;
+ if (proxy.find("@") != std::string::npos)
+ {
+ std::string udata = getuntil("@", proxy);
+ proxy = getafter("@", proxy);
+
+ std::string username = getuntil(":", proxy);
+ std::string password = getafter(":", proxy);
+ std::string adata = username + ":" + password;
+
+ authorization = "Authorization: Basic " + base64_encode(reinterpret_cast(adata.c_str()), adata.size())+"\r\n";
+ }
+
+ unsigned short port = ssl ? 443 : 80;
+ if (proxy.find(":") != std::string::npos)
+ {
+ port = static_cast(watoi(getafter(":", proxy)));
+ proxy = getuntil(":", proxy);
+ }
+
+ IPipe* cs;
+ if (ssl)
+ cs = Server->ConnectSslStream(proxy, port, 10000);
+ else
+ cs = Server->ConnectStream(proxy, port, 10000);
+
+ if (cs == NULL)
+ return cs;
+
+ std::string connect_data = "CONNECT " + selected_server_settings.hostname + ":" + convert(selected_server_settings.port) + " HTTP/1.1\r\n"
+ +"Host: "+ selected_server_settings.hostname+"\r\n"
+ + authorization+"\r\n";
+
+ if (!cs->Write(connect_data))
+ {
+ Server->destroy(cs);
+ return NULL;
+ }
+
+ char buf[512];
+ int64 starttime = Server->getTimeMS();
+ int state = 0;
+ std::string state_data;
+ int http_code = 0;
+ do
+ {
+ size_t rc = cs->Read(buf, sizeof(buf), 10000);
+
+ if (rc == 0)
+ {
+ break;
+ }
+
+ for (size_t i = 0; i < rc; ++i)
+ {
+ char ch = buf[i];
+ switch (state)
+ {
+ case 0:
+ {
+ if (ch == ' ')
+ {
+ if (strlower(state_data) != "http/1.1"
+ && strlower(state_data) != "http/1.0")
+ {
+ Server->Log("Unknown HTTP protocol: " + state_data, LL_ERROR);
+ Server->destroy(cs);
+ return NULL;
+ }
+ state_data.clear();
+ ++state;
+ }
+ else
+ state_data += ch;
+ }break;
+ case 1:
+ {
+ if (ch == ' ')
+ {
+ http_code = watoi(state_data);
+ state_data.clear();
+ ++state;
+ }
+ else
+ state_data += ch;
+ }break;
+ case 2:
+ {
+ if (ch == '\n')
+ {
+ if (http_code != 200)
+ {
+ Server->Log("HTTP proxy returned error code " + convert(http_code) + " message \"" + state_data + "\"", LL_ERROR);
+ Server->destroy(cs);
+ return NULL;
+ }
+ state_data.clear();
+ ++state;
+ }
+ else if (ch != '\r')
+ {
+ state_data += ch;
+ }
+ } break;
+ case 3:
+ {
+ if (ch == '\n')
+ {
+ if (i + 1 < rc)
+ {
+ tcpstack.AddData(&buf[i + 1], rc - i - 1);
+ }
+ return cs;
+ }
+ else if (ch != '\r')
+ {
+ ++state;
+ }
+ } break;
+ case 4:
+ {
+ if (ch == '\n')
+ {
+ state = 3;
+ }
+ }break;
+ default:
+ {
+ assert(false);
+ Server->destroy(cs);
+ return NULL;
+ }break;
+ }
+ }
+
+ } while (Server->getTimeMS() - starttime < 10000);
+
+ Server->Log("Timeout connecting via http proxy");
+
+ Server->destroy(cs);
+ return NULL;
+ }
+
+ return Server->ConnectStream(selected_server_settings.hostname,
+ selected_server_settings.port, 10000);
+}
diff --git a/urbackupclient/InternetClient.h b/urbackupclient/InternetClient.h
index ce28549c..80993b1c 100644
--- a/urbackupclient/InternetClient.h
+++ b/urbackupclient/InternetClient.h
@@ -13,9 +13,16 @@ class ICustomClient;
class IScopedLock;
class ICondition;
+struct SServerConnectionSettings
+{
+ std::string hostname;
+ std::string proxy;
+ unsigned short port;
+};
+
struct SServerSettings
{
- std::vector > servers;
+ std::vector servers;
size_t selected_server;
std::string clientname;
std::string authkey;
@@ -58,6 +65,8 @@ public:
static void setStatusMsg(const std::string& msg);
+ static IPipe* connect(const SServerConnectionSettings& selected_settings, CTCPStack& tcpstack);
+
private:
static IMutex *mutex;
@@ -77,7 +86,8 @@ private:
class InternetClientThread : public IThread
{
public:
- InternetClientThread(IPipe *cs, const SServerSettings &server_settings);
+ InternetClientThread(IPipe *cs, const SServerSettings &server_settings, CTCPStack* tcpstack);
+ ~InternetClientThread();
void operator()(void);
char *getReply(CTCPStack *tcpstack, IPipe *pipe, size_t &replysize, unsigned int timeoutms);
@@ -88,5 +98,6 @@ private:
std::string generateRandomBinaryAuthKey(void);
void printInfo( IPipe * pipe );
IPipe *cs;
+ CTCPStack* tcpstack;
SServerSettings server_settings;
};
diff --git a/urbackupcommon/settingslist.cpp b/urbackupcommon/settingslist.cpp
index 71e567bb..c0fdbb3f 100644
--- a/urbackupcommon/settingslist.cpp
+++ b/urbackupcommon/settingslist.cpp
@@ -56,6 +56,7 @@ std::vector getSettingsList(void)
ret.push_back("image_letters");
ret.push_back("internet_server");
ret.push_back("internet_server_port");
+ ret.push_back("internet_server_proxy");
ret.push_back("internet_authkey");
ret.push_back("internet_speed");
ret.push_back("local_speed");
@@ -146,6 +147,7 @@ std::vector getGlobalizedSettingsList(void)
std::vector ret;
ret.push_back("internet_server");
ret.push_back("internet_server_port");
+ ret.push_back("internet_server_proxy");
ret.push_back("server_url");
return ret;
}
@@ -172,6 +174,7 @@ std::vector getGlobalSettingsList(void)
ret.push_back("backup_database");
ret.push_back("internet_server");
ret.push_back("internet_server_port");
+ ret.push_back("internet_server_proxy");
ret.push_back("global_local_speed");
ret.push_back("global_internet_speed");
ret.push_back("use_tmpfiles");
diff --git a/urbackupserver/server_settings.cpp b/urbackupserver/server_settings.cpp
index 36327522..b751e81d 100644
--- a/urbackupserver/server_settings.cpp
+++ b/urbackupserver/server_settings.cpp
@@ -369,6 +369,7 @@ void ServerSettings::readSettingsDefault(ISettingsReader* settings_default,
settings->image_letters=settings_default->getValue("image_letters", "C");
settings->backup_database=(settings_global->getValue("backup_database", "true")=="true");
settings->internet_server_port=(unsigned short)(atoi(settings_global->getValue("internet_server_port", "55415").c_str()));
+ settings->internet_server_proxy = settings_global->getValue("internet_server_proxy", "");
settings->client_set_settings=false;
settings->internet_server= settings_global->getValue("internet_server", "");
settings->internet_image_backups=(settings_default->getValue("internet_image_backups", "false")=="true");
diff --git a/urbackupserver/server_settings.h b/urbackupserver/server_settings.h
index 9067e942..e4c97128 100644
--- a/urbackupserver/server_settings.h
+++ b/urbackupserver/server_settings.h
@@ -79,6 +79,7 @@ struct SSettings
std::string internet_server;
bool client_set_settings;
unsigned short internet_server_port;
+ std::string internet_server_proxy;
std::string internet_authkey;
bool internet_full_file_backups;
bool internet_image_backups;
diff --git a/urbackupserver/serverinterface/add_client.cpp b/urbackupserver/serverinterface/add_client.cpp
index 0626f36f..dd92358d 100644
--- a/urbackupserver/serverinterface/add_client.cpp
+++ b/urbackupserver/serverinterface/add_client.cpp
@@ -29,6 +29,7 @@ ACTION_IMPL(add_client)
ret.set("new_authkey", new_authkey);
ret.set("internet_server", s->internet_server);
ret.set("internet_server_port", s->internet_server_port);
+ ret.set("internet_server_proxy", s->internet_server_proxy);
ret.set("added_new_client", true);
}
else
diff --git a/urbackupserver/serverinterface/download_client.cpp b/urbackupserver/serverinterface/download_client.cpp
index b6efd393..bd7f83c9 100644
--- a/urbackupserver/serverinterface/download_client.cpp
+++ b/urbackupserver/serverinterface/download_client.cpp
@@ -55,6 +55,7 @@ namespace
ret+="internet_mode_enabled="+convert(settingsptr->internet_mode_enabled)+"\r\n";
ret+="internet_server="+settingsptr->internet_server+"\r\n";
ret+="internet_server_port="+convert(settingsptr->internet_server_port)+"\r\n";
+ ret += "internet_server_proxy=" + settingsptr->internet_server_proxy + "\r\n";
ret+="internet_authkey="+(authkey.empty() ? settingsptr->internet_authkey : authkey ) +"\r\n";
if(!clientname.empty())
{
diff --git a/urbackupserver/serverinterface/settings.cpp b/urbackupserver/serverinterface/settings.cpp
index fc5707ec..50d2eef8 100644
--- a/urbackupserver/serverinterface/settings.cpp
+++ b/urbackupserver/serverinterface/settings.cpp
@@ -234,6 +234,7 @@ void getGeneralSettings(JSON::Object& obj, IDatabase *db, ServerSettings &settin
SET_SETTING(backup_database);
SET_SETTING(internet_server);
SET_SETTING(internet_server_port);
+ SET_SETTING(internet_server_proxy);
SET_SETTING(global_local_speed);
SET_SETTING(global_soft_fs_quota);
SET_SETTING(global_internet_speed);
diff --git a/urbackupserver/www/templates/settings_inv_row.htm b/urbackupserver/www/templates/settings_inv_row.htm
index 7ff78d62..8d043d4a 100644
--- a/urbackupserver/www/templates/settings_inv_row.htm
+++ b/urbackupserver/www/templates/settings_inv_row.htm
@@ -464,6 +464,12 @@
+
{/global_settings}
{?main_client}
{^global_settings}