Skip to content

Commit b88e6c6

Browse files
committed
Revert "Revert "Add ip_matcher module, converted from nodejs""
This reverts commit 2fc934b.
1 parent 2fc934b commit b88e6c6

13 files changed

Lines changed: 1940 additions & 0 deletions

File tree

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
"""
2+
Based on https://github.com/demskie/netparser
3+
MIT License - Copyright (c) 2019 alex
4+
"""
5+
6+
from .shared import parse_base_network, sort_networks, summarize_sorted_networks
7+
from .sort import binary_search_for_insertion_index
8+
9+
10+
class IPMatcher:
11+
def __init__(self, networks=None):
12+
self.sorted = []
13+
if networks is not None:
14+
subnets = []
15+
for s in networks:
16+
net = parse_base_network(s, False)
17+
if net and net.is_valid():
18+
subnets.append(net)
19+
sort_networks(subnets)
20+
self.sorted = summarize_sorted_networks(subnets)
21+
22+
def has(self, network):
23+
"""
24+
Checks if the given IP address is in the list of networks.
25+
"""
26+
net = parse_base_network(network, False)
27+
if not net or not net.is_valid():
28+
return False
29+
idx = binary_search_for_insertion_index(net, self.sorted)
30+
if idx < len(self.sorted) and self.sorted[idx].contains(net):
31+
return True
32+
if idx - 1 >= 0 and self.sorted[idx - 1].contains(net):
33+
return True
34+
return False
35+
36+
def add(self, network):
37+
net = parse_base_network(network, False)
38+
if not net or not net.is_valid():
39+
return self
40+
idx = binary_search_for_insertion_index(net, self.sorted)
41+
if idx < len(self.sorted) and self.sorted[idx].compare(net) == 0:
42+
return self
43+
self.sorted.insert(idx, net)
44+
return self
Lines changed: 119 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,119 @@
1+
"""
2+
Based on https://github.com/demskie/netparser
3+
MIT License - Copyright (c) 2019 alex
4+
"""
5+
6+
from aikido_zen.helpers.ip_matcher import parse
7+
8+
# Constants
9+
BEFORE = -1
10+
EQUALS = 0
11+
AFTER = 1
12+
13+
14+
class Address:
15+
def __init__(self, address=None):
16+
self.arr = []
17+
if address:
18+
net = parse.network(address)
19+
if net:
20+
self.arr = net["bytes"]
21+
22+
def bytes(self):
23+
return self.arr if self.arr else []
24+
25+
def set_bytes(self, bytes_data):
26+
if len(bytes_data) == 4 or len(bytes_data) == 16:
27+
self.arr = bytes_data
28+
else:
29+
self.arr = []
30+
return self
31+
32+
def destroy(self):
33+
if self.is_valid():
34+
self.arr = []
35+
return self
36+
37+
def is_valid(self):
38+
return len(self.arr) > 0
39+
40+
def is_ipv4(self):
41+
return len(self.arr) == 4
42+
43+
def is_ipv6(self):
44+
return len(self.arr) == 16
45+
46+
def duplicate(self):
47+
return Address().set_bytes(self.arr.copy())
48+
49+
def equals(self, address):
50+
return self.compare(address) == EQUALS
51+
52+
def compare(self, address):
53+
if not self.is_valid() or not address.is_valid():
54+
return None
55+
if self == address:
56+
return EQUALS
57+
if len(self.arr) < len(address.arr):
58+
return BEFORE
59+
if len(self.arr) > len(address.arr):
60+
return AFTER
61+
62+
for i in range(len(self.arr)):
63+
if self.arr[i] < address.arr[i]:
64+
return BEFORE
65+
if self.arr[i] > address.arr[i]:
66+
return AFTER
67+
68+
return EQUALS
69+
70+
def apply_subnet_mask(self, cidr):
71+
if not self.is_valid():
72+
return self
73+
mask_bits = len(self.arr) * 8 - cidr
74+
for i in range(len(self.arr) - 1, -1, -1):
75+
mask = max(0, min(mask_bits, 8))
76+
if mask == 0:
77+
return self
78+
self.arr[i] &= ~((1 << mask) - 1)
79+
mask_bits -= 8
80+
return self
81+
82+
def is_base_address(self, cidr):
83+
if not self.is_valid() or cidr < 0 or cidr > len(self.arr) * 8:
84+
return False
85+
if cidr == len(self.arr) * 8:
86+
return True
87+
mask_bits = len(self.arr) * 8 - cidr
88+
for i in range(len(self.arr) - 1, -1, -1):
89+
mask = max(0, min(mask_bits, 8))
90+
if mask == 0:
91+
return True
92+
if self.arr[i] != (self.arr[i] & ~((1 << mask) - 1)):
93+
return False
94+
mask_bits -= 8
95+
return True
96+
97+
def increase(self, cidr):
98+
if self.is_valid():
99+
self.offset_address(cidr, True)
100+
else:
101+
self.destroy()
102+
return self
103+
104+
def offset_address(self, cidr, forwards, throw_errors=False):
105+
target_byte = (cidr - 1) // 8
106+
if self.is_valid() and 0 <= target_byte < len(self.arr):
107+
increment = 2 ** (8 - (cidr - target_byte * 8))
108+
self.arr[target_byte] += increment * (1 if forwards else -1)
109+
if target_byte >= 0:
110+
if self.arr[target_byte] < 0:
111+
self.arr[target_byte] = 256 + (self.arr[target_byte] % 256)
112+
self.offset_address(target_byte * 8, forwards, throw_errors)
113+
elif self.arr[target_byte] > 255:
114+
self.arr[target_byte] %= 256
115+
self.offset_address(target_byte * 8, forwards, throw_errors)
116+
else:
117+
self.destroy()
118+
else:
119+
self.destroy()
Lines changed: 49 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,49 @@
1+
import pytest
2+
from .address import Address
3+
4+
5+
def test_is_ipv4():
6+
assert Address("1.2.3.5").is_ipv4() == True
7+
assert Address("::1").is_ipv4() == False
8+
9+
10+
def test_is_ipv6():
11+
assert Address("::1").is_ipv6() == True
12+
13+
14+
def test_compare():
15+
assert Address("1.2.3.4").compare(Address("1.2.3.5")) == -1
16+
assert Address("1.2.3.4").compare(Address("1.2.3.4")) == 0
17+
assert Address("1.2.3.4").compare(Address("1.2.3.4").duplicate()) == 0
18+
19+
# edge cases
20+
assert Address().compare(Address()) is None
21+
assert Address("1.2.3.4").compare(Address()) is None
22+
23+
24+
def test_bytes_method():
25+
assert Address("1.2.3.4").bytes() == [1, 2, 3, 4]
26+
assert Address().bytes() == []
27+
28+
29+
def test_set_bytes_method():
30+
assert Address("1.2.3.4").set_bytes([3]).bytes() == []
31+
32+
33+
def test_equals_method():
34+
assert Address("1.2.3.4").equals(Address("1.2.3.4")) == True
35+
assert Address("1.2.3.4").equals(Address("1.2.3.5")) == False
36+
37+
38+
def test_apply_subnet_mask():
39+
assert Address().apply_subnet_mask(0) is not None
40+
41+
42+
def test_increase_method():
43+
assert Address().increase(0).bytes() == []
44+
assert Address().bytes() == []
45+
46+
47+
def test_self_comparison():
48+
a = Address("3.4.5.6")
49+
assert a.compare(a) == 0

0 commit comments

Comments
 (0)