From 89d44ffbbb0c0ed8346f511b554ef00fb5d00730 Mon Sep 17 00:00:00 2001 From: DL6ER Date: Thu, 13 Apr 2017 07:18:55 +0200 Subject: [PATCH] Add magic byte to each struct entry that is checked before accessing the variable - if the check fails, print a warning to the log file --- FTL.h | 1 + gc.c | 24 +++++++++---------- parser.c | 69 +++++++++++++++++++++++++++++++++++++----------------- request.c | 32 ++++++++++++------------- routines.h | 2 +- 5 files changed, 78 insertions(+), 50 deletions(-) diff --git a/FTL.h b/FTL.h index c0b14675..a9cd54aa 100644 --- a/FTL.h +++ b/FTL.h @@ -141,6 +141,7 @@ typedef struct { } clientsDataStruct; typedef struct { + unsigned char magic; int count; int blockedcount; char *domain; diff --git a/gc.c b/gc.c index 1618c7a6..632c8ab1 100644 --- a/gc.c +++ b/gc.c @@ -32,20 +32,20 @@ void *GC_thread(void *val) int invalidated = 0; for(i=0; i < counters.queries; i++) { - validate_access("queries", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("queries", i, true, __LINE__, __FUNCTION__, __FILE__); if(queries[i].timestamp < mintime && queries[i].valid) { // Adjust total counters and total over time data // We cannot edit counters.queries directly as it is used // as max ID for the queries[] struct counters.invalidqueries++; - validate_access("overTime", queries[i].timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", queries[i].timeidx, true, __LINE__, __FUNCTION__, __FILE__); overTime[queries[i].timeidx].total--; // Adjust client and domain counters - validate_access("clients", queries[i].clientID, __LINE__, __FUNCTION__, __FILE__); + validate_access("clients", queries[i].clientID, true, __LINE__, __FUNCTION__, __FILE__); clients[queries[i].clientID].count--; - validate_access("domains", queries[i].domainID, __LINE__, __FUNCTION__, __FILE__); + validate_access("domains", queries[i].domainID, true, __LINE__, __FUNCTION__, __FILE__); domains[queries[i].domainID].count--; // Change other counters according to status of this query @@ -56,14 +56,14 @@ void *GC_thread(void *val) break; case 1: counters.blocked--; - validate_access("overTime", queries[i].timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", queries[i].timeidx, true, __LINE__, __FUNCTION__, __FILE__); overTime[queries[i].timeidx].blocked--; - validate_access("domains", queries[i].domainID, __LINE__, __FUNCTION__, __FILE__); + validate_access("domains", queries[i].domainID, true, __LINE__, __FUNCTION__, __FILE__); domains[queries[i].domainID].blockedcount--; break; case 2: counters.forwardedqueries--; - validate_access("forwarded", queries[i].forwardID, __LINE__, __FUNCTION__, __FILE__); + validate_access("forwarded", queries[i].forwardID, true, __LINE__, __FUNCTION__, __FILE__); forwarded[queries[i].forwardID].count--; break; case 3: @@ -71,7 +71,7 @@ void *GC_thread(void *val) break; case 4: counters.wildcardblocked--; - validate_access("overTime", queries[i].timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", queries[i].timeidx, true, __LINE__, __FUNCTION__, __FILE__); overTime[queries[i].timeidx].blocked--; break; default: @@ -83,12 +83,12 @@ void *GC_thread(void *val) { case 1: counters.IPv4--; - validate_access("overTime", queries[i].timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", queries[i].timeidx, true, __LINE__, __FUNCTION__, __FILE__); overTime[queries[i].timeidx].querytypedata[0]--; break; case 2: counters.IPv6--; - validate_access("overTime", queries[i].timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", queries[i].timeidx, true, __LINE__, __FUNCTION__, __FILE__); overTime[queries[i].timeidx].querytypedata[1]--; break; default: @@ -100,8 +100,8 @@ void *GC_thread(void *val) int j; for(j = 0; j < overTime[queries[i].timeidx].forwardnum; j++) { - validate_access("forwarded", j, __LINE__, __FUNCTION__, __FILE__); - validate_access("overTime", queries[i].timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("forwarded", j, true, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", queries[i].timeidx, true, __LINE__, __FUNCTION__, __FILE__); forwarded[j].count -= overTime[queries[i].timeidx].forwarddata[j]; validate_access_oTfd(queries[i].timeidx, j, __LINE__, __FUNCTION__, __FILE__); diff --git a/parser.c b/parser.c index 5507f6c8..016371dc 100644 --- a/parser.c +++ b/parser.c @@ -9,6 +9,7 @@ * Please see LICENSE file for your rights under this license. */ #include "FTL.h" +#define MAGICBYTE 0x57 char *resolveHostname(char *addr); void extracttimestamp(char *readbuffer, int *querytimestamp, int *overTimetimestamp); @@ -198,7 +199,7 @@ void process_pihole_log(int file) bool found = false; for(i=0; i < counters.overTime; i++) { - validate_access("overTime", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", i, true, __LINE__, __FUNCTION__, __FILE__); if(overTime[i].timestamp == overTimetimestamp) { found = true; @@ -210,7 +211,9 @@ void process_pihole_log(int file) { memory_check(OVERTIME); timeidx = counters.overTime; - validate_access("overTime", timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", timeidx, false, __LINE__, __FUNCTION__, __FILE__); + // Set magic byte + overTime[timeidx].magic = MAGICBYTE; overTime[timeidx].timestamp = overTimetimestamp; overTime[timeidx].total = 0; overTime[timeidx].blocked = 0; @@ -266,14 +269,14 @@ void process_pihole_log(int file) { type = 1; counters.IPv4++; - validate_access("overTime", timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", timeidx, true, __LINE__, __FUNCTION__, __FILE__); overTime[timeidx].querytypedata[0]++; } else if(strstr(readbuffer,"query[AAAA]") != NULL) { type = 2; counters.IPv6++; - validate_access("overTime", timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", timeidx, true, __LINE__, __FUNCTION__, __FILE__); overTime[timeidx].querytypedata[1]++; } @@ -363,7 +366,9 @@ void process_pihole_log(int file) // // Debug output if(debug) logg("New domain: %s (%i/%i)", domain, domainID, counters.domains_MAX); - validate_access("domains", domainID, __LINE__, __FUNCTION__, __FILE__); + validate_access("domains", domainID, false, __LINE__, __FUNCTION__, __FILE__); + // Set magic byte + domains[domainID].magic = MAGICBYTE; // Set its counter to 1 domains[domainID].count = 1; // Set blocked counter to zero @@ -395,7 +400,9 @@ void process_pihole_log(int file) else logg("New client: %s (%i/%i)", client, clientID, counters.clients_MAX); - validate_access("clients", clientID, __LINE__, __FUNCTION__, __FILE__); + validate_access("clients", clientID, false, __LINE__, __FUNCTION__, __FILE__); + // Set magic byte + clients[clientID].magic = MAGICBYTE; // Set its counter to 1 clients[clientID].count = 1; // Store client IP @@ -412,7 +419,8 @@ void process_pihole_log(int file) } // Save everything - validate_access("queries", queryID, __LINE__, __FUNCTION__, __FILE__); + validate_access("queries", queryID, false, __LINE__, __FUNCTION__, __FILE__); + queries[queryID].magic = MAGICBYTE; queries[queryID].timestamp = querytimestamp; queries[queryID].type = type; queries[queryID].status = status; @@ -426,7 +434,7 @@ void process_pihole_log(int file) counters.queries++; // Update overTime data - validate_access("overTime", timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", timeidx, true, __LINE__, __FUNCTION__, __FILE__); overTime[timeidx].total++; // Decide what to increment depending on status @@ -437,9 +445,9 @@ void process_pihole_log(int file) break; case 1: counters.blocked++; - validate_access("overTime", timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", timeidx, true, __LINE__, __FUNCTION__, __FILE__); overTime[timeidx].blocked++; - validate_access("domains", domainID, __LINE__, __FUNCTION__, __FILE__); + validate_access("domains", domainID, true, __LINE__, __FUNCTION__, __FILE__); domains[domainID].blockedcount++; break; case 2: @@ -450,9 +458,9 @@ void process_pihole_log(int file) break; case 4: counters.wildcardblocked++; - validate_access("overTime", timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", timeidx, true, __LINE__, __FUNCTION__, __FILE__); overTime[timeidx].blocked++; - validate_access("domains", domainID, __LINE__, __FUNCTION__, __FILE__); + validate_access("domains", domainID, true, __LINE__, __FUNCTION__, __FILE__); domains[domainID].wildcard = true; break; default: @@ -480,7 +488,7 @@ void process_pihole_log(int file) bool found = false; for(i=0; i < counters.overTime; i++) { - validate_access("overTime", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", i, true, __LINE__, __FUNCTION__, __FILE__); if(overTime[i].timestamp == overTimetimestamp) { found = true; @@ -492,7 +500,8 @@ void process_pihole_log(int file) { memory_check(OVERTIME); timeidx = counters.overTime; - validate_access("overTime", timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", timeidx, false, __LINE__, __FUNCTION__, __FILE__); + overTime[timeidx].magic = MAGICBYTE; overTime[timeidx].timestamp = overTimetimestamp; overTime[timeidx].total = 0; overTime[timeidx].blocked = 0; @@ -504,7 +513,7 @@ void process_pihole_log(int file) } // Determine if there is enough space for saving the current // forwardID in the overTime data structure -allocate space otherwise - validate_access("overTime", timeidx, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", timeidx, true, __LINE__, __FUNCTION__, __FILE__); if(overTime[timeidx].forwardnum <= forwardID) { // Reallocate more space for forwarddata @@ -601,7 +610,7 @@ int detectStatus(char *domain) char part[strlen(domain)],partbuffer[strlen(domain)]; for(i=0; i < counters.wildcarddomains; i++) { - validate_access("wildcarddomains", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("wildcarddomains", i, false, __LINE__, __FUNCTION__, __FILE__); if(strcmp(wildcarddomains[i], domain) == 0) { // Exact match with wildcard domain @@ -700,7 +709,7 @@ int getforwardID(char * str) // Go through already knows forward servers and see if we used one of those for(i=0; i < counters.forwarded; i++) { - validate_access("forwarded", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("forwarded", i, true, __LINE__, __FUNCTION__, __FILE__); if(strcmp(forwarded[i].ip, forward) == 0) { forwardID = i; @@ -723,7 +732,9 @@ int getforwardID(char * str) else logg("New forward server: %s (%i/%u)", forward, forwardID, counters.forwarded_MAX); - validate_access("forwarded", forwardID, __LINE__, __FUNCTION__, __FILE__); + validate_access("forwarded", forwardID, false, __LINE__, __FUNCTION__, __FILE__); + // Set magic byte + forwarded[forwardID].magic = MAGICBYTE; // Set its counter to 1 forwarded[forwardID].count = 1; // Save IP @@ -750,7 +761,7 @@ int findDomain(char *domain) int i; for(i=0; i < counters.domains; i++) { - validate_access("domains", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("domains", i, true, __LINE__, __FUNCTION__, __FILE__); // Quick test: Does the domain start with the same character? if(domains[i].domain[0] != domain[0]) continue; @@ -771,7 +782,7 @@ int findClient(char *client) int i; for(i=0; i < counters.clients; i++) { - validate_access("clients", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("clients", i, true, __LINE__, __FUNCTION__, __FILE__); // Quick test: Does the clients IP start with the same character? if(clients[i].ip[0] != client[0]) continue; @@ -787,7 +798,7 @@ int findClient(char *client) return -1; } -void validate_access(const char * name, int pos, int line, const char * function, const char * file) +void validate_access(const char * name, int pos, bool testmagic, int line, const char * function, const char * file) { int limit = 0; if(name[0] == 'c') limit = counters.clients_MAX; @@ -796,11 +807,27 @@ void validate_access(const char * name, int pos, int line, const char * function else if(name[0] == 'o') limit = counters.overTime_MAX; else if(name[0] == 'f') limit = counters.forwarded_MAX; else if(name[0] == 'w') limit = counters.wildcarddomains; + else { logg("Validator error"); killed = 1; } if(pos >= limit || pos < 0) { logg("FATAL ERROR: Trying to access %s[%i], but maximum is %i", name, pos, limit); logg(" found in %s() (line %i) in %s", function, line, file); } + if(testmagic) + { + unsigned char magic = 0x00; + if(name[0] == 'c') magic = clients[pos].magic; + else if(name[0] == 'd') magic = domains[pos].magic; + else if(name[0] == 'q') magic = queries[pos].magic; + else if(name[0] == 'o') magic = overTime[pos].magic; + else if(name[0] == 'f') magic = forwarded[pos].magic; + else { logg("Validator error"); killed = 1; } + if(magic != MAGICBYTE) + { + logg("FATAL ERROR: Trying to access %s[%i], but magic byte is %x", name, pos, magic); + logg(" found in %s() (line %i) in %s", function, line, file); + } + } } void validate_access_oTfd(int timeidx, int pos, int line, const char * function, const char * file) diff --git a/request.c b/request.c index 48aed6a3..4252139d 100644 --- a/request.c +++ b/request.c @@ -221,7 +221,7 @@ void getOverTime(int *sock) bool sendit = false; for(i=0; i < counters.overTime; i++) { - validate_access("overTime", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", i, true, __LINE__, __FUNCTION__, __FILE__); if((overTime[i].total > 0 || overTime[i].blocked > 0) && !sendit) { sendit = true; @@ -255,7 +255,7 @@ void getTopDomains(char *client_message, int *sock) for(i=0; i < counters.domains; i++) { - validate_access("domains", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("domains", i, true, __LINE__, __FUNCTION__, __FILE__); temparray[i][0] = i; if(blocked) temparray[i][1] = domains[i].blockedcount; @@ -299,7 +299,7 @@ void getTopDomains(char *client_message, int *sock) { // Get sorted indices int j = temparray[counters.domains-i-1][0]; - validate_access("domains", j, __LINE__, __FUNCTION__, __FILE__); + validate_access("domains", j, true, __LINE__, __FUNCTION__, __FILE__); // Skip this domain if there is a filter on it if(excludedomains != NULL) @@ -346,7 +346,7 @@ void getTopClients(char *client_message, int *sock) for(i=0; i < counters.clients; i++) { - validate_access("clients", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("clients", i, true, __LINE__, __FUNCTION__, __FILE__); temparray[i][0] = i; temparray[i][1] = clients[i].count; } @@ -364,7 +364,7 @@ void getTopClients(char *client_message, int *sock) { // Get sorted indices int j = temparray[counters.clients-i-1][0]; - validate_access("clients", j, __LINE__, __FUNCTION__, __FILE__); + validate_access("clients", j, true, __LINE__, __FUNCTION__, __FILE__); // Skip this client if there is a filter on it if(excludeclients != NULL) @@ -396,7 +396,7 @@ void getForwardDestinations(int *sock) int i, temparray[counters.forwarded][2]; for(i=0; i < counters.forwarded; i++) { - validate_access("forwarded", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("forwarded", i, true, __LINE__, __FUNCTION__, __FILE__); temparray[i][0] = i; temparray[i][1] = forwarded[i].count; } @@ -408,7 +408,7 @@ void getForwardDestinations(int *sock) { // Get sorted indices int j = temparray[counters.forwarded-i-1][0]; - validate_access("forwarded", j, __LINE__, __FUNCTION__, __FILE__); + validate_access("forwarded", j, true, __LINE__, __FUNCTION__, __FILE__); if(forwarded[j].count > 0) { sprintf(server_message,"%i %i %s %s\n",i,forwarded[j].count,forwarded[j].ip,forwarded[j].name); @@ -427,7 +427,7 @@ void getForwardNames(int *sock) for(i=0; i < counters.forwarded; i++) { - validate_access("forwarded", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("forwarded", i, true, __LINE__, __FUNCTION__, __FILE__); // Get sorted indices sprintf(server_message,"%i %i %s %s\n",i,forwarded[i].count,forwarded[i].ip,forwarded[i].name); swrite(server_message, *sock); @@ -556,12 +556,12 @@ void getAllQueries(char *client_message, int *sock) int i; for(i=ibeg; i < counters.queries; i++) { - validate_access("queries", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("queries", i, true, __LINE__, __FUNCTION__, __FILE__); // Check if this query has been removed due to garbage collection if(!queries[i].valid) continue; - validate_access("domains", queries[i].domainID, __LINE__, __FUNCTION__, __FILE__); - validate_access("clients", queries[i].clientID, __LINE__, __FUNCTION__, __FILE__); + validate_access("domains", queries[i].domainID, true, __LINE__, __FUNCTION__, __FILE__); + validate_access("clients", queries[i].clientID, true, __LINE__, __FUNCTION__, __FILE__); char type[5]; if(queries[i].type == 1) @@ -648,7 +648,7 @@ void getRecentBlocked(char *client_message, int *sock) int found = 0; for(i = counters.queries - 1; i > 0 ; i--) { - validate_access("queries", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("queries", i, true, __LINE__, __FUNCTION__, __FILE__); // Check if this query has been removed due to garbage collection if(!queries[i].valid) continue; @@ -701,7 +701,7 @@ void getForwardDestinationsOverTime(int *sock) int i, sendit = -1; for(i = 0; i < counters.overTime; i++) { - validate_access("overTime", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", i, true, __LINE__, __FUNCTION__, __FILE__); if((overTime[i].total > 0 || overTime[i].blocked > 0)) { sendit = i; @@ -712,7 +712,7 @@ void getForwardDestinationsOverTime(int *sock) { for(i = sendit; i < counters.overTime; i++) { - validate_access("overTime", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", i, true, __LINE__, __FUNCTION__, __FILE__); sprintf(server_message, "%i", overTime[i].timestamp); int j; @@ -752,7 +752,7 @@ void getQueryTypesOverTime(int *sock) int i, sendit = -1; for(i = 0; i < counters.overTime; i++) { - validate_access("overTime", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", i, true, __LINE__, __FUNCTION__, __FILE__); if((overTime[i].total > 0 || overTime[i].blocked > 0)) { sendit = i; @@ -763,7 +763,7 @@ void getQueryTypesOverTime(int *sock) { for(i = sendit; i < counters.overTime; i++) { - validate_access("overTime", i, __LINE__, __FUNCTION__, __FILE__); + validate_access("overTime", i, true, __LINE__, __FUNCTION__, __FILE__); sprintf(server_message, "%i %i %i\n", overTime[i].timestamp,overTime[i].querytypedata[0],overTime[i].querytypedata[1]); swrite(server_message, *sock); } diff --git a/routines.h b/routines.h index 076de50b..abdadbdd 100644 --- a/routines.h +++ b/routines.h @@ -28,7 +28,7 @@ void open_pihole_log(void); void handle_signals(void); void process_pihole_log(int file); void *pihole_log_thread(void *val); -void validate_access(const char * name, int pos, int line, const char * function, const char * file); +void validate_access(const char * name, int pos, bool testmagic, int line, const char * function, const char * file); void validate_access_oTfd(int timeidx, int pos, int line, const char * function, const char * file); void pihole_log_flushed(bool message);