Files
digFTP/src/main.cpp

567 lines
16 KiB
C++

#include <iostream>
#include <string>
#include <unistd.h>
#include <cstring>
#include <future>
#include <thread>
#include <chrono>
#include <stdio.h>
#include <stdlib.h>
#include <signal.h>
#include <csignal>
#include <mutex>
#include <memory>
#include <system_error>
#include <sys/ioctl.h>
#include <sys/poll.h>
#include <sys/socket.h>
#include <sys/time.h>
#include <netinet/in.h>
#include <errno.h>
#include <fcntl.h>
#include "main.h"
#include "client.cpp"
#include "util.h"
using namespace std::chrono_literals;
std::mutex client_mutex;
struct pollfd fds[MAX_CLIENTS];
struct ftpconn {
Client* client;
std::thread* thread;
bool close = false;
} fdc[MAX_CLIENTS];
void runClient(struct ftpconn* cfd);
void initializePlugins();
void shutdown(int signum);
int main(int argc , char *argv[]) {
std::cout <<
std::format(
"{} {} Maintained by Worlio LLC 2024-2025\n"
"This program comes with ABSOLUTELY NO WARRANTY.\n"
"This is free software, and you are welcome to redistribute it under certain conditions.\n\n",
APPNAME,
VERSION
);
// SIGNALS
signal(SIGPIPE, SIG_IGN);
std::signal(SIGINT, shutdown);
config = new ConfigFile(concatPath(std::string(CONFIG_DIR), "ftp.conf"));
server_name = config->getValue("core", "server_name", "%a %v");
motd = config->getValue("core", "motd", "cmd:cowsay -r Welcome %u!");
sscanf(
config->getValue("net", "listen", "127.0.0.1:21").c_str(),
"%u.%u.%u.%u:%u",
&server_address[0],
&server_address[1],
&server_address[2],
&server_address[3],
&server_port
);
logger = new Logger();
std::string mainLogFile = config->getValue("logging", "all", "");
logger->openFileOnLevel(LOGLEVEL_MAX, mainLogFile.c_str());
logger->openFileOnLevel(LOGLEVEL_DEBUG, config->getValue("logging", "debug", mainLogFile).c_str());
logger->openFileOnLevel(LOGLEVEL_INFO, config->getValue("logging", "info", mainLogFile).c_str());
logger->openFileOnLevel(LOGLEVEL_WARNING, config->getValue("logging", "warning", mainLogFile).c_str());
logger->openFileOnLevel(LOGLEVEL_ERROR, config->getValue("logging", "error", mainLogFile).c_str());
logger->openFileOnLevel(LOGLEVEL_CRITICAL, config->getValue("logging", "critical", mainLogFile).c_str());
std::string console_loglevel_string = config->getValue("logging", "console", "all");
if (console_loglevel_string == "all")
logger->setConsoleLevel(LOGLEVEL_MAX);
else if (console_loglevel_string == "debug")
logger->setConsoleLevel(LOGLEVEL_DEBUG);
else if (console_loglevel_string == "info")
logger->setConsoleLevel(LOGLEVEL_INFO);
else if (console_loglevel_string == "warning")
logger->setConsoleLevel(LOGLEVEL_WARNING);
else if (console_loglevel_string == "error")
logger->setConsoleLevel(LOGLEVEL_ERROR);
else if (console_loglevel_string == "critical")
logger->setConsoleLevel(LOGLEVEL_CRITICAL);
else if (console_loglevel_string == "none")
logger->setConsoleLevel(LOGLEVEL_MIN);
else {
logger->setConsoleLevel(LOGLEVEL_MIN);
logger->print(LOGLEVEL_ERROR, "Could not determine console log type.");
}
bool ssl_enable = config->getBool("net", "ssl", true);
bool ssl_flags = SSL_OP_NO_TICKET;
if (ssl_enable) {
std::string cert_file = config->getValue("ssl", "certificate", "cert.pem");
if (cert_file[0] != '/')
cert_file = concatPath(std::string(CONFIG_DIR), cert_file);
logger->print(LOGLEVEL_INFO, "Using certificate file: {}", cert_file);
std::string key_file = config->getValue("ssl", "private_key", "key.pem");
if (key_file[0] != '/')
key_file = concatPath(std::string(CONFIG_DIR), key_file);
logger->print(LOGLEVEL_INFO, "Using private key file: {}", key_file);
if (!SSLManager::getInstance().initialize(cert_file, key_file)) {
logger->print(LOGLEVEL_CRITICAL, "Failed to initialize SSL");
return 1;
}
// set configured ciphers
SSLManager::getInstance().setCiphers(config->getValue("ssl", "ciphers", "ALL:!ADH:!LOW:!EXP:!MD5:@STRENGTH").c_str());
// handle configured flags
if (!config->getBool("ssl", "ssl_v2", false))
ssl_flags+=SSL_OP_NO_SSLv2;
if (!config->getBool("ssl", "ssl_v3", false))
ssl_flags+=SSL_OP_NO_SSLv3;
if (!config->getBool("ssl", "tls_v1_0", false))
ssl_flags+=SSL_OP_NO_TLSv1;
if (!config->getBool("ssl", "tls_v1_1", false))
ssl_flags+=SSL_OP_NO_TLSv1_1;
if (!config->getBool("ssl", "tls_v1_2", true))
ssl_flags+=SSL_OP_NO_TLSv1_2;
if (!config->getBool("ssl", "tls_v1_3", true))
ssl_flags+=SSL_OP_NO_TLSv1_3;
if (!(ssl_flags & SSL_OP_NO_SSLv2) || !(ssl_flags & SSL_OP_NO_SSLv3))
global_features.push_back("AUTH SSL");
if (!(ssl_flags & SSL_OP_NO_TLSv1) || !(ssl_flags & SSL_OP_NO_TLSv1_1) ||
!(ssl_flags & SSL_OP_NO_TLSv1_2) || !(ssl_flags & SSL_OP_NO_TLSv1_3))
global_features.push_back("AUTH TLS");
global_features.push_back("PBSZ");
global_features.push_back("PROT");
if ((ssl_flags & SSL_OP_NO_SSLv2) && (ssl_flags & SSL_OP_NO_SSLv3) &&
(ssl_flags & SSL_OP_NO_TLSv1) && (ssl_flags & SSL_OP_NO_TLSv1_1) &&
(ssl_flags & SSL_OP_NO_TLSv1_2) && (ssl_flags & SSL_OP_NO_TLSv1_3))
logger->print(LOGLEVEL_WARNING, "All SSL/TLS protocols disabled. You're a mad man!");
if (!config->getBool("ssl", "compression", true))
ssl_flags+=SSL_OP_NO_COMPRESSION;
if (config->getBool("ssl", "prefer_server_ciphers", true))
ssl_flags+=SSL_OP_CIPHER_SERVER_PREFERENCE;
SSLManager::getInstance().setFlags(ssl_flags);
}
initializePlugins();
int opt = 1,
master_socket = -1,
newsock = -1,
nfds = 1,
current_size = 0,
src = 0;
runServer = true;
runCompression = false;
struct sockaddr_in ctrl_address;
if ((master_socket = socket(AF_INET, SOCK_STREAM, 0)) < 0) {
logger->print(LOGLEVEL_CRITICAL, "Failed creating socket");
return master_socket;
}
if ((src = setsockopt(master_socket, SOL_SOCKET, SO_REUSEADDR, (char *)&opt, sizeof(opt))) < 0) {
logger->print(LOGLEVEL_CRITICAL, "Unable to configure socket");
close(master_socket);
return src;
}
if ((src = ioctl(master_socket, FIONBIO, (char *)&opt)) < 0) {
logger->print(LOGLEVEL_CRITICAL, "Unable to read socket");
close(master_socket);
return src;
}
ctrl_address.sin_family = AF_INET;
ctrl_address.sin_addr.s_addr = INADDR_ANY;
ctrl_address.sin_port = htons(server_port);
if ((src = bind(master_socket, (struct sockaddr *)&ctrl_address, sizeof(ctrl_address))) < 0) {
logger->print(
LOGLEVEL_CRITICAL,
"Bind to {}.{}.{}.{}:{} failed",
static_cast<unsigned int>(server_address[0]),
static_cast<unsigned int>(server_address[1]),
static_cast<unsigned int>(server_address[2]),
static_cast<unsigned int>(server_address[3]),
server_port
);
close(master_socket);
return src;
}
if ((src = listen(master_socket, 3)) < 0) {
logger->print(LOGLEVEL_CRITICAL, "Unable to listen to socket");
close(master_socket);
return src;
}
memset(fds, 0, sizeof(fds));
memset(fdc, 0, sizeof(fdc));
fds[0].fd = master_socket;
fds[0].events = POLLIN;
logger->print(LOGLEVEL_INFO, "Server started.");
while (runServer) {
int pc = poll(fds, nfds, -1);
if (pc < 0) {
if (errno == EINTR) continue;
logger->print(LOGLEVEL_CRITICAL, "Connection poll faced a fatal error");
break;
}
if (pc == 0) continue;
current_size = nfds;
for (int i = 0; i < current_size; i++) {
if (fds[i].revents == 0)
continue;
// Handle poll errors properly without skipping cleanup
if (fds[i].revents != POLLIN) {
if (fds[i].fd != master_socket) {
logger->print(LOGLEVEL_ERROR, "net: poll error on fd {}", fds[i].fd);
fdc[i].close = true;
}
}
if (fds[i].fd == master_socket && fds[i].revents == POLLIN) {
// Handle new connections
do {
struct sockaddr_in client_name;
socklen_t client_len = sizeof(client_name);
newsock = accept(master_socket, (struct sockaddr*)&client_name, &client_len);
if (newsock < 0) {
if (errno != EWOULDBLOCK) {
logger->print(LOGLEVEL_ERROR, "net: accept() failed: {}", strerror(errno));
runServer = false;
}
break;
}
// Find first available slot
int slot = -1;
for (int j = 1; j < MAX_CLIENTS; j++) {
if (fds[j].fd <= 0) { // Changed from < 0 to <= 0
slot = j;
break;
}
}
if (slot < 0) {
logger->print(LOGLEVEL_ERROR, "No free slots available");
close(newsock);
continue;
}
// Set non-blocking mode
int flags = fcntl(newsock, F_GETFL, 0);
fcntl(newsock, F_SETFL, flags | O_NONBLOCK);
char client_addr[INET_ADDRSTRLEN];
inet_ntop(AF_INET, &client_name.sin_addr, client_addr, sizeof(client_addr));
logger->print(
LOGLEVEL_DEBUG,
"client {} (ip: {}) accepted in slot {}",
newsock,
client_addr,
slot
);
fds[slot].fd = newsock;
fds[slot].events = POLLIN;
fds[slot].revents = 0;
fdc[slot].client = new Client(newsock);
for (const auto &o : client_options)
fdc[slot].client->addFeature(o, "");
fdc[slot].thread = new std::thread(runClient, &fdc[slot]);
fdc[slot].close = false;
if (slot >= nfds) {
nfds = slot + 1;
}
} while (newsock != -1);
}
// Handle cleanup for any connections marked for closing
if (fds[i].fd != master_socket && (fdc[i].close || fds[i].revents != POLLIN)) {
int fd = fds[i].fd;
logger->print(LOGLEVEL_DEBUG, "cleaning up client {} from slot {}", fd, i);
// Close socket
if (fd > 0) {
shutdown(fd, SHUT_RDWR);
close(fd);
}
// Clean up thread
if (fdc[i].thread) {
if (fdc[i].thread->joinable()) {
fdc[i].thread->join();
}
delete fdc[i].thread;
}
// Clean up client
delete fdc[i].client;
// Reset slot
fds[i].fd = -1;
fds[i].events = 0;
fds[i].revents = 0;
memset(&fdc[i], 0, sizeof(struct ftpconn));
// Recalculate nfds if needed
if (i == nfds - 1) {
for (int j = nfds - 1; j >= 0; j--) {
if (fds[j].fd != -1) {
nfds = j + 1;
break;
}
}
}
}
}
}
logger->print(LOGLEVEL_INFO, "Server closing...");
close(master_socket);
for (int i = 0; i < current_size; i++) {
if (fds[i].fd != master_socket) {
int fd = fds[i].fd;
logger->print(LOGLEVEL_DEBUG, "disconnecting client {} from slot {}", fd, i);
if (fd > 0) {
shutdown(fd, SHUT_RDWR);
close(fd);
}
if (fdc[i].thread) {
if (fdc[i].thread->joinable()) {
fdc[i].thread->join();
}
delete fdc[i].thread;
}
// Clean up client
delete fdc[i].client;
// Reset slot
fds[i].fd = -1;
fds[i].events = 0;
fds[i].revents = 0;
memset(&fdc[i], 0, sizeof(struct ftpconn));
if (i == nfds - 1) {
for (int j = nfds - 1; j >= 0; j--) {
if (fds[j].fd != -1) {
nfds = j + 1;
break;
}
}
}
}
}
logger->close();
return 0;
}
void initializePlugins() {
std::string plugin_dir = config->getValue("core", "plugin_path", PLUGIN_DIR);
auto& auth_manager = PluginManager<Auth>::getInstance();
auto& filer_manager = PluginManager<Filer>::getInstance();
auth_manager.setLogger(logger);
filer_manager.setLogger(logger);
// Try loading all plugins into both managers
for (const auto& entry : std::filesystem::directory_iterator(plugin_dir)) {
if (entry.path().extension() == ".so") {
logger->print(LOGLEVEL_DEBUG, "Loading plugin: {}", entry.path().string());
auth_manager.loadPlugin(entry.path().string());
filer_manager.loadPlugin(entry.path().string());
}
}
// Initialize auth
std::string auth_type = config->getValue("engines", "auth", "pam");
auth = auth_manager.createPlugin(auth_type, config->get(auth_type)->get());
if (!auth) {
logger->print(LOGLEVEL_CRITICAL, "Failed to create auth engine: {}", auth_type);
exit(1);
}
// Initialize filer
std::string filer_type = config->getValue("engines", "filer", "local");
default_filer_name = filer_type;
default_filer_factory = filer_manager.getFactory(filer_type);
if (!default_filer_factory) {
logger->print(LOGLEVEL_CRITICAL, "Failed to get filer factory for type {}", filer_type);
exit(1);
}
}
void runClient(struct ftpconn* cfd) {
if (!cfd) {
logger->print(LOGLEVEL_ERROR, "Invalid connection handle");
return;
}
std::unique_lock<std::mutex> lock(client_mutex);
if (!cfd->client) {
logger->print(LOGLEVEL_ERROR, "Invalid client handle");
return;
}
int client_sock = cfd->client->control_sock;
Client* client = cfd->client;
lock.unlock();
char inbuf[BUFFERSIZE];
logger->print(LOGLEVEL_DEBUG, "client {} initialized", client_sock);
while (!cfd->close) {
memset(inbuf, 0, BUFFERSIZE);
if (fcntl(client_sock, F_GETFD) < 0) {
logger->print(LOGLEVEL_DEBUG, "closed client {} socket", client_sock);
break;
}
struct timeval tv;
tv.tv_sec = 60;
tv.tv_usec = 0;
fd_set readfds, writefds;
FD_ZERO(&readfds);
FD_ZERO(&writefds);
FD_SET(client_sock, &readfds);
// Add socket to writefds if SSL wants to write
if (client->isSecure() && client->getSSL()) {
FD_SET(client_sock, &writefds);
}
int select_result = select(client_sock + 1, &readfds, &writefds, NULL, &tv);
if (select_result < 0) {
if (errno == EINTR) continue;
logger->print(LOGLEVEL_ERROR, "client {} experienced select fail: %s", client_sock, strerror(errno));
break;
}
if (select_result == 0) {
logger->print(LOGLEVEL_INFO, "client {} timeout", client_sock);
break;
}
int rc;
if (client->isSecure() && client->getSSL()) {
if (!client->isHandshakeComplete()) {
// Continue SSL handshake
ERR_clear_error(); // Clear any previous errors
int ret = SSL_accept(client->getSSL());
if (ret <= 0) {
int ssl_err = SSL_get_error(client->getSSL(), ret);
if (ssl_err == SSL_ERROR_WANT_READ || ssl_err == SSL_ERROR_WANT_WRITE) {
continue; // Need more data for handshake
}
unsigned long err = ERR_get_error();
char err_buf[256];
ERR_error_string_n(err, err_buf, sizeof(err_buf));
logger->print(LOGLEVEL_ERROR, "client {} SSL handshake failed with error: {} ({})",
client_sock, ssl_err, err_buf);
break;
}
client->setHandshakeComplete(true);
logger->print(LOGLEVEL_DEBUG, "client {} SSL handshake completed", client_sock);
continue;
} else {
// Normal SSL read after handshake
ERR_clear_error(); // Clear any previous errors
rc = SSL_read(client->getSSL(), inbuf, sizeof(inbuf) - 1);
if (rc <= 0) {
int ssl_err = SSL_get_error(client->getSSL(), rc);
if (ssl_err == SSL_ERROR_WANT_READ || ssl_err == SSL_ERROR_WANT_WRITE) {
continue;
}
if (ssl_err == SSL_ERROR_SYSCALL) {
unsigned long err = ERR_get_error();
if (err == 0 && rc == 0) {
logger->print(LOGLEVEL_DEBUG, "client {} SSL connection closed", client_sock);
} else {
char err_buf[256];
ERR_error_string_n(err, err_buf, sizeof(err_buf));
logger->print(LOGLEVEL_ERROR, "client {} experienced SSL_read syscall error: {}",
client_sock, err_buf);
}
} else {
unsigned long err = ERR_get_error();
char err_buf[256];
ERR_error_string_n(err, err_buf, sizeof(err_buf));
logger->print(LOGLEVEL_ERROR, "client {} experienced SSL_read error: {} ({})",
client_sock, ssl_err, err_buf);
}
break;
}
}
} else {
rc = recv(client_sock, inbuf, sizeof(inbuf) - 1, 0);
if (rc <= 0) {
if (rc == 0) {
logger->print(LOGLEVEL_DEBUG, "client {} disconnected", client_sock);
} else {
logger->print(LOGLEVEL_ERROR, "recieve from client {} failed: {}", client_sock, strerror(errno));
}
break;
}
}
inbuf[rc] = '\0';
if (rc >= 2 && inbuf[rc-2] == '\r' && inbuf[rc-1] == '\n') {
rc -= 2;
inbuf[rc] = '\0';
}
std::string input(inbuf, rc);
logger->print(LOGLEVEL_DEBUG, "recieved from client {}: {}", client_sock, input);
std::string::size_type space_pos = input.find(" ");
std::string cmd = space_pos != std::string::npos ?
toUpper(input.substr(0, space_pos)) : toUpper(input);
std::string args = space_pos != std::string::npos ?
input.substr(space_pos + 1) : "";
lock.lock();
if (!cfd->client) {
lock.unlock();
break;
}
int revc = client->receive(cmd, args);
lock.unlock();
if (revc != 0) break;
}
// Mark for cleanup
logger->print(LOGLEVEL_DEBUG, "client {} thread ending", client_sock);
cfd->close = true;
}
void shutdown(int signum) {
std::cout << "\n";
runServer = false;
}