net.py 3.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114
  1. import json, os, sys
  2. from Crypto.Cipher import AES
  3. # import lib.config
  4. # should be 32 bytes hex
  5. # https://pycryptodome.readthedocs.io/en/latest/src/cipher/aes.html
  6. # KEY = bytes.fromhex(lib.config.get(
  7. # "shared_secret",
  8. # "87b9b70e722d20c046c8dba8d0add1f16307fec33debffec9d001fd20dbca3ee"
  9. # ))
  10. class Channel:
  11. def __init__(self, reader, writer):
  12. self.reader = reader
  13. self.writer = writer
  14. async def readline(self):
  15. if not (line := await self.reader.readline()):
  16. self.writer.close()
  17. return None
  18. # Strip the newline
  19. return line[:-1].decode()
  20. async def receive(self):
  21. if (plaintext := await self.readline()) is None:
  22. return None
  23. # if (ciphertext := await self.readline()) is None:
  24. # return None
  25. # if (tag := await self.readline()) is None:
  26. # return None
  27. # nonce = bytes.fromhex(nonce)
  28. # ciphertext = bytes.fromhex(ciphertext)
  29. # tag = bytes.fromhex(tag)
  30. #print(f"{nonce.hex()}")
  31. #print(f"{ciphertext.hex()}")
  32. #print(f"{tag.hex()}")
  33. #print()
  34. # Decrypt
  35. # cipher = AES.new(KEY, AES.MODE_EAX, nonce=nonce)
  36. # plaintext = cipher.decrypt(ciphertext)
  37. # try:
  38. # cipher.verify(tag)
  39. # except ValueError:
  40. # print("error: key incorrect or message corrupted", file=sys.stderr)
  41. # return None
  42. message = plaintext
  43. response = json.loads(message)
  44. return response
  45. async def send(self, obj):
  46. message = json.dumps(obj)
  47. data = message.encode()
  48. # Encrypt
  49. # cipher = AES.new(KEY, AES.MODE_EAX)
  50. # nonce = cipher.nonce
  51. # ciphertext, tag = cipher.encrypt_and_digest(data)
  52. #print(f"{nonce.hex()}")
  53. #print(f"{ciphertext.hex()}")
  54. #print(f"{tag.hex()}")
  55. #print()
  56. # Encode as hex strings since the bytes might contain new lines
  57. # nonce = nonce.hex().encode()
  58. # ciphertext = ciphertext.hex().encode()
  59. # tag = tag.hex().encode()
  60. # self.writer.write(nonce + b"\n")
  61. # self.writer.write(ciphertext + b"\n")
  62. # self.writer.write(tag + b"\n")
  63. self.writer.write(data + b"\n")
  64. await self.writer.drain()
  65. async def _test_client():
  66. await asyncio.sleep(1)
  67. reader, writer = await asyncio.open_connection("127.0.0.1", 7643)
  68. channel = Channel(reader, writer)
  69. request = {
  70. "foo": "bar"
  71. }
  72. await channel.send(request)
  73. response = await channel.receive()
  74. print(f"Client: {response}")
  75. async def _test_server(reader, writer):
  76. channel = Channel(reader, writer)
  77. request = await channel.receive()
  78. print(f"Server: {request}")
  79. response = {
  80. "abc": "xyz"
  81. }
  82. await channel.send(response)
  83. async def _test_channel():
  84. server = await asyncio.start_server(_test_server, "127.0.0.1", 7643)
  85. task1 = asyncio.create_task(_test_client())
  86. async with server:
  87. task2 = asyncio.create_task(server.serve_forever())
  88. await asyncio.sleep(3)
  89. await task1
  90. task2.cancel()
  91. if __name__ == "__main__":
  92. import asyncio
  93. # run send and recv testes
  94. asyncio.run(_test_channel())