From 92b7b7beffdba8b99b5b21543d60ad5708f92d1e Mon Sep 17 00:00:00 2001 From: Cloud Wu Date: Mon, 16 Mar 2015 20:01:22 +0800 Subject: [PATCH 1/3] use table for multi header key --- lualib/http/httpd.lua | 9 ++++++++- lualib/http/internal.lua | 10 +++++++++- 2 files changed, 17 insertions(+), 2 deletions(-) diff --git a/lualib/http/httpd.lua b/lualib/http/httpd.lua index c19d4bae..0131db0a 100644 --- a/lualib/http/httpd.lua +++ b/lualib/http/httpd.lua @@ -2,6 +2,7 @@ local internal = require "http.internal" local table = table local string = string +local type = type local httpd = {} @@ -113,7 +114,13 @@ local function writeall(writefunc, statuscode, bodyfunc, header) writefunc(statusline) if header then for k,v in pairs(header) do - writefunc(string.format("%s: %s\r\n", k,v)) + if type(v) == "table" then + for _,v in ipairs(v) do + writefunc(string.format("%s: %s\r\n", k,v)) + end + else + writefunc(string.format("%s: %s\r\n", k,v)) + end end end local t = type(bodyfunc) diff --git a/lualib/http/internal.lua b/lualib/http/internal.lua index 0e83bc72..5bb35752 100644 --- a/lualib/http/internal.lua +++ b/lualib/http/internal.lua @@ -1,3 +1,6 @@ +local table = table +local type = type + local M = {} local LIMIT = 8192 @@ -83,7 +86,12 @@ function M.parseheader(lines, from, header) end name = name:lower() if header[name] then - header[name] = header[name] .. ", " .. value + local v = header[name] + if type(v) == "table" then + table.insert(v, value) + else + header[name] = { v , value } + end else header[name] = value end From a9b1993686ae0cb46336352dae2493ff01e2a11a Mon Sep 17 00:00:00 2001 From: Cloud Wu Date: Sun, 29 Mar 2015 14:41:28 +0800 Subject: [PATCH 2/3] Add async dns query --- lualib/dns.lua | 290 +++++++++++++++++++++++++++++++++++++++++++++++ test/testdns.lua | 11 ++ 2 files changed, 301 insertions(+) create mode 100644 lualib/dns.lua create mode 100644 test/testdns.lua diff --git a/lualib/dns.lua b/lualib/dns.lua new file mode 100644 index 00000000..4636bf6d --- /dev/null +++ b/lualib/dns.lua @@ -0,0 +1,290 @@ +--[[ + lua dns resolver library + See https://github.com/xjdrew/levent/blob/master/levent/dns.lua for more detail + +-- resource record type: +-- TYPE value and meaning +-- A 1 a host address +-- NS 2 an authoritative name server +-- MD 3 a mail destination (Obsolete - use MX) +-- MF 4 a mail forwarder (Obsolete - use MX) +-- CNAME 5 the canonical name for an alias +-- SOA 6 marks the start of a zone of authority +-- MB 7 a mailbox domain name (EXPERIMENTAL) +-- MG 8 a mail group member (EXPERIMENTAL) +-- MR 9 a mail rename domain name (EXPERIMENTAL) +-- NULL 10 a null RR (EXPERIMENTAL) +-- WKS 11 a well known service description +-- PTR 12 a domain name pointer +-- HINFO 13 host information +-- MINFO 14 mailbox or mail list information +-- MX 15 mail exchange +-- TXT 16 text strings +-- AAAA 28 a ipv6 host address +-- only appear in the question section: +-- AXFR 252 A request for a transfer of an entire zone +-- MAILB 253 A request for mailbox-related records (MB, MG or MR) +-- MAILA 254 A request for mail agent RRs (Obsolete - see MX) +-- * 255 A request for all records +-- +-- resource recode class: +-- IN 1 the Internet +-- CS 2 the CSNET class (Obsolete - used only for examples in some obsolete RFCs) +-- CH 3 the CHAOS class +-- HS 4 Hesiod [Dyer 87] +-- only appear in the question section: +-- * 255 any class +-- ]] + +--[[ +-- struct header { +-- uint16_t tid # identifier assigned by the program that generates any kind of query. +-- uint16_t flags # flags +-- uint16_t qdcount # the number of entries in the question section. +-- uint16_t ancount # the number of resource records in the answer section. +-- uint16_t nscount # the number of name server resource records in the authority records section. +-- uint16_t arcount # the number of resource records in the additional records section. +-- } +-- +-- request body: +-- struct request { +-- string name +-- uint16_t atype +-- uint16_t class +-- } +-- +-- response body: +-- struct response { +-- string name +-- uint16_t atype +-- uint16_t class +-- uint16_t ttl +-- uint16_t rdlength +-- string rdata +-- } +--]] + +local skynet = require "skynet" +local socket = require "socket" + +local MAX_DOMAIN_LEN = 1024 +local MAX_LABEL_LEN = 63 +local MAX_PACKET_LEN = 2048 +local DNS_HEADER_LEN = 12 + +local QTYPE = { + A = 1, + CNAME = 5, + AAAA = 28, +} + +local QCLASS = { + IN = 1, +} + +local dns = {} + +local function verify_domain_name(name) + if #name > MAX_DOMAIN_LEN then + return false + end + if not name:match("^[%l%d-%.]+$") then + return false + end + for w in name:gmatch("([%w-]+)%.?") do + if #w > MAX_LABEL_LEN then + return false + end + end + return true +end + +local next_tid = 1 +local function gen_tid() + local tid = next_tid + next_tid = next_tid + 1 + return tid +end + +local function pack_header(t) + return string.pack(">HHHHHH", + t.tid, t.flags, t.qdcount, t.ancount or 0, t.nscount or 0, t.arcount or 0) +end + +local function pack_question(name, qtype, qclass) + local labels = {} + for w in name:gmatch("([%w-]+)%.?") do + table.insert(labels, string.pack("s1",w)) + end + table.insert(labels, '\0') + table.insert(labels, string.pack(">HH", qtype, qclass)) + return table.concat(labels) +end + +local function unpack_header(chunk) + local tid, flags, qdcount, ancount, nscount, arcount, left = string.unpack(">HHHHHH", chunk) + return { + tid = tid, + flags = flags, + qdcount = qdcount, + ancount = ancount, + nscount = nscount, + arcount = arcount + }, left +end + +-- unpack a resource name +local function unpack_name(chunk, left) + local t = {} + local jump_pointer + local tag, offset, label + while true do + tag, left = string.unpack("B", chunk, left) + if tag & 0xc0 == 0xc0 then + -- pointer + offset,left = string.unpack(">H", chunk, left - 1) + offset = offset & 0x3fff + if not jump_pointer then + jump_pointer = left + end + -- offset is base 0, need to plus 1 + left = offset + 1 + elseif tag == 0 then + break + else + label, left = string.unpack("s1", chunk, left - 1) + t[#t+1] = label + end + end + return table.concat(t, "."), jump_pointer or left +end + +local function unpack_question(chunk, left) + local name, left = unpack_name(chunk, left) + local atype, class, left = string.unpack(">HH", chunk, left) + return { + name = name, + atype = atype, + class = class + }, left +end + +local function unpack_answer(chunk, left) + local name, left = unpack_name(chunk, left) + local atype, class, ttl, rdata, left = string.unpack(">HHI4s2", chunk, left) + return { + name = name, + atype = atype, + class = class, + ttl = ttl, + rdata = rdata + },left +end + +local function unpack_rdata(qtype, chunk) + if qtype == QTYPE.A then + local a,b,c,d = string.unpack("BBBB", chunk) + return string.format("%d.%d.%d.%d", a,b,c,d) + elseif qtype == QTYPE.AAAA then + local a,b,c,d,e,f,g,h = string.unpack(">HHHHHHHH", chunk) + return string.format("%x:%x:%x:%x:%x:%x:%x:%x", a, b, c, d, e, f, g, h) + else + error("Error qtype " .. qtype) + end +end + +local dns_server +local request_pool = {} + +local function resolve(content) + if #content < DNS_HEADER_LEN then + -- drop + skynet.error("Recv an invalid package when dns query") + return + end + local answer_header,left = unpack_header(content) + -- verify answer + assert(answer_header.qdcount == 1, "malformed packet") + + local resp = request_pool[answer_header.tid] + if not resp then + skynet.error("Recv an invalid tid when dns query") + return + end + + local question,left = unpack_question(content, left) + if question.name ~= resp.name then + skynet.error("Recv an invalid name when dns query") + return + end + + local ttl + local answer + local answers = {} + for i=1, answer_header.ancount do + answer, left = unpack_answer(content, left) + -- only extract qtype address + if answer.atype == resp.qtype then + local ip = unpack_rdata(resp.qtype, answer.rdata) + ttl = ttl and math.min(ttl, answer.ttl) or answer.ttl + answers[#answers+1] = ip + end + end + + if #answers > 0 then + resp.answers = answers + end + + skynet.wakeup(resp.co) +end + +function dns.server(server, port) + if not server then + local f = assert(io.open "/etc/resolv.conf") + for line in f:lines() do + server = line:match("%s*nameserver%s+([^%s]+)") + if server then + break + end + end + assert(server, "Can't get nameserver") + end + dns_server = socket.udp(function(data, sz, from) + resolve(skynet.tostring(data,sz)) + end) + socket.udp_connect(dns_server, server, port or 53) + return server +end + +local function suspend(tid, name, qtype) + local req = { + name = name, + tid = tid, + qtype = qtype, + time = skynet.now(), -- for timeout + co = coroutine.running(), + } + request_pool[tid] = req + skynet.wait() + local answers = request_pool[tid].answers + request_pool[tid] = nil + assert(answers, "no ip") + return answers[1], answers +end + +function dns.resolve(name, ipv6) + local qtype = ipv6 and QTYPE.AAAA or QTYPE.A + local name = name:lower() + assert(verify_domain_name(name) , "illegal name") + local question_header = { + tid = gen_tid(), + flags = 0x100, -- flags: 00000001 00000000, set RD + qdcount = 1, + } + local req = pack_header(question_header) .. pack_question(name, qtype, QCLASS.IN) + assert(dns_server, "Call dns.server fist") + socket.write(dns_server, req) + return suspend(question_header.tid, name, qtype) +end + +return dns diff --git a/test/testdns.lua b/test/testdns.lua new file mode 100644 index 00000000..1d767ce0 --- /dev/null +++ b/test/testdns.lua @@ -0,0 +1,11 @@ +local skynet = require "skynet" +local dns = require "dns" + +skynet.start(function() + print("nameserver:", dns.server()) -- set nameserver + -- you can specify the server like dns.server("8.8.4.4", 53) + local ip, ips = dns.resolve "github.com" + for k,v in ipairs(ips) do + print("github.com",v) + end +end) From 36ed8679b869c17f84e1a46db4f702470385c5ea Mon Sep 17 00:00:00 2001 From: Cloud Wu Date: Mon, 30 Mar 2015 09:29:22 +0800 Subject: [PATCH 3/3] you can set host in httpc header --- lualib/http/httpc.lua | 7 ++++++- test/testhttp.lua | 14 +++++++++++--- 2 files changed, 17 insertions(+), 4 deletions(-) diff --git a/lualib/http/httpc.lua b/lualib/http/httpc.lua index a73d5710..823d33b5 100644 --- a/lualib/http/httpc.lua +++ b/lualib/http/httpc.lua @@ -14,10 +14,15 @@ local function request(fd, method, host, url, recvheader, header, content) for k,v in pairs(header) do header_content = string.format("%s%s:%s\r\n", header_content, k, v) end + if header.host then + host = "" + end + else + host = string.format("host:%s\r\n",host) end if content then - local data = string.format("%s %s HTTP/1.1\r\nhost:%s\r\ncontent-length:%d\r\n%s\r\n%s", method, url, host, #content, header_content, content) + local data = string.format("%s %s HTTP/1.1\r\n%scontent-length:%d\r\n%s\r\n%s", method, url, host, #content, header_content, content) write(data) else local request_header = string.format("%s %s HTTP/1.1\r\nhost:%s\r\ncontent-length:0\r\n%s\r\n", method, url, host, header_content) diff --git a/test/testhttp.lua b/test/testhttp.lua index 9fd0804c..cba1488a 100644 --- a/test/testhttp.lua +++ b/test/testhttp.lua @@ -1,16 +1,24 @@ local skynet = require "skynet" local httpc = require "http.httpc" +local dns = require "dns" skynet.start(function() print("GET baidu.com") - local header = {} - local status, body = httpc.get("baidu.com", "/", header) + local respheader = {} + local status, body = httpc.get("baidu.com", "/", respheader) print("[header] =====>") - for k,v in pairs(header) do + for k,v in pairs(respheader) do print(k,v) end print("[body] =====>", status) print(body) + local respheader = {} + dns.server() + local ip = dns.resolve "baidu.com" + print(string.format("GET %s (baidu.com)", ip)) + local status, body = httpc.get("baidu.com", "/", respheader, { host = "baidu.com" }) + print(status) + skynet.exit() end)