#!/usr/bin/python

'''
igrep, Internet sniffer.

(c) 2015 Jan ONDREJ (SAL) <ondrejj(at)salstar.sk>

 This program is free software; you can redistribute it and/or modify
 it under the terms of the GNU General Public License as published by
 the Free Software Foundation; either version 2 of the License, or
 (at your option) any later version.

Usage: igrep.py [options]

  --help | -h		This help.
  -i interface		Set sniffing interface to "interface".
  -e expr		Search for "expr".
'''

import sys, re, time, pcapy, getopt, string, unicodedata

class unknown:
  type = 'unknown'
  def __init__(self, payload):
      self.payload = payload
  def data_offset(self):
      return 0
  def data(self):
      return self.payload[self.data_offset():]

class tcp(unknown):
  type = 'tcp'
  def src(self):
      return (ord(self.payload[0])<<8)+ord(self.payload[1])
  def dst(self): 
      return (ord(self.payload[2])<<8)+ord(self.payload[3])
  def data_offset(self):
      return (ord(self.payload[12]) & 0xf0) >> 2
  def len(self):
      return len(self.payload)-self.data_offset() 
  def is_fin(self):
      return (ord(self.payload[13])&1)==1

class udp(tcp):
  type = 'udp'
  def data_offset(self):
      #print len(self.payload), (ord(self.payload[4])<<8)+ord(self.payload[5])
      return 8

class icmp(tcp):
  # https://en.wikipedia.org/wiki/Internet_Control_Message_Protocol
  type = 'icmp'
  types = {
    0:	"Echo reply",
    3:	"Destination Unreachable",
    8:	"Echo Request"
  }
  codes = {
    3: {
      0:	"Destination network unreachable",
      1:	"Destination host unreachable",
      2:	"Destination protocol unreachable",
      3:	"Destination port unreachable",
      4:	"Fragmentation required, and DF flag set",
      5:	"Source route failed",
      6:	"Destination network unknown",
      7:	"Destination host unknown",
      8:	"Source host isolated",
      9:	"Network administratively prohibited",
      10:	"Host administratively prohibited",
      11:	"Network unreachable for TOS",
      12:	"Host unreachable for TOS",
      13:	"Communication administratively prohibited",
      14:	"Host Precedence Violation",
      15:	"Precedence cutoff in effect"
    }
  }
  def __init__(self, payload):
      self.payload = payload
      self.protocol_type = ord(payload[0])
      self.protocol_code = ord(payload[1])
      if self.protocol_type==3:
        self.tcp = ipv4(payload[8:])
      else:
        self.tcp = None
  def data_offset(self):
      return 8
  def as_string(self):
      return "%s: %s" % (
        self.types.get(self.protocol_type, self.protocol_type),
        self.codes.get(self.protocol_code, self.protocol_code)
      )

class ipv4:
  protocols = {
    1: icmp,
    6: tcp,
    17: udp
  }
  def __init__(self, payload):
      self.payload = payload
  def ihl(self):
      return (ord(self.payload[0]) & 0x0f) << 2
  def protocol(self):
      return ord(self.payload[9])
  def src(self):
      return [ord(x) for x in self.payload[12:16]]
  def dst(self):
      return [ord(x) for x in self.payload[16:20]]
  def as_string(self, addr):
      return '.'.join([str(x) for x in addr])
  def ssrc(self):
      return self.as_string(self.src())
  def sdst(self):
      return self.as_string(self.dst())
  def data(self):
      protocol = self.protocols.get(self.protocol(), unknown)
      return protocol(self.payload[self.ihl():])

def si(key, unit='B', delimeter=' '):
    fix = ['', 'k', 'M', 'G', 'T', 'P', 'E', 'Z', 'Y']
    fix_len = len(fix)-1
    counter = 0
    while (key >= 1024) and (fix_len>counter):
      key /= 1024
      counter += 1
      if len(fix)==(counter-1):
        break
    return "%4d%s%1s%s" % (key, delimeter, fix[counter], unit)

control_chars = [
  unichr(x) for x in range(0,32) + range(127,256)
  if unichr(x) not in u' \r\n\t'
]
control_char_re = re.compile('[%s]' % re.escape(''.join(control_chars)))

def remove_control_chars(s):
    return control_char_re.sub('.', s)

def format_time(ts):
    return time.strftime(
      "%Y-%m-%d %H:%M:%S" + (".%06d" % ts[1]),
      time.localtime(ts[0])
    )

if __name__ == "__main__":
  opts, args = getopt.gnu_getopt(sys.argv[1:], 'hi:e:', [
    'help', 'interface=', 'expression='
  ])
  opts = dict(opts)
  if opts.has_key("--help") or opts.has_key("-h"):
    print __doc__
    sys.exit(0)

  if opts.has_key("--interface"):
    interface = opts['--interface']
  elif opts.has_key("-i"):
    interface = opts['-i']
  else:
    interface = pcapy.findalldevs()[0]

  cap = pcapy.open_live(interface, 1500, 0, 100)
  print "Capture interface %s: %s/%s" \
        % (interface, cap.getnet(), cap.getmask())

  if opts.has_key("--expression"):
    grep = re.compile(opts['--expression']).search
  elif opts.has_key("-e"):
    grep = re.compile(opts['-e']).search
  else:
    def grep(x):
        return True

  if len(args)>1:
    ip_filter = ' '.join(args)
    print "Filter: %s" % ip_filter
    cap.setfilter(ip_filter)

  try:
    while True:
      try:
        header, payload = cap.next()
      except pcapy.PcapError:
        continue
      dg = ipv4(payload[14:])
      proto = dg.data()
      if proto.type=="tcp" or proto.type=="udp":
        data = proto.data()
        if grep(data):
          print "%s: %s:%d -> %s:%d" \
                % (format_time(header.getts()),
                   dg.ssrc(), proto.src(), dg.sdst(), proto.dst())
          print remove_control_chars(data)
      elif proto.type=="icmp":
        print "%s: %s -> %s" \
              % (format_time(header.getts()), dg.ssrc(), dg.sdst())
        print "ICMP:", proto.as_string()
        if proto.tcp:
          e = proto.tcp.data()
          print "TCP: %s:%d -> %s:%d" \
                % (proto.tcp.ssrc(), e.src(), proto.tcp.sdst(), e.dst())
  except KeyboardInterrupt:
    pass
  except Exception, e:
    raise
