From: Michal 'vorner' Vaner Date: Wed, 7 Mar 2012 10:20:08 +0000 (+0100) Subject: Merge #1601 X-Git-Tag: trac2351_base~226^2~116^2~41^2~50 X-Git-Url: http://git.ipfire.org/gitweb.cgi?a=commitdiff_plain;h=d7fb4b7244e60a034e21873d5bd32f7148ccd973;p=thirdparty%2Fkea.git Merge #1601 Conflicts: src/bin/auth/auth_srv.cc src/bin/auth/tests/auth_srv_unittest.cc --- d7fb4b7244e60a034e21873d5bd32f7148ccd973 diff --cc src/bin/auth/auth_srv.cc index 23751f2771,3026f96e11..524d581fec --- a/src/bin/auth/auth_srv.cc +++ b/src/bin/auth/auth_srv.cc @@@ -291,32 -290,30 +290,32 @@@ makeErrorMessage(Message& message, Outp // If this is an error to a query or notify, we should also copy the // question section. if (opcode == Opcode::QUERY() || opcode == Opcode::NOTIFY()) { - questions.assign(message->beginQuestion(), message->endQuestion()); + questions.assign(message.beginQuestion(), message.endQuestion()); } - message->clear(Message::RENDER); - message->setQid(qid); - message->setOpcode(opcode); - message->setHeaderFlag(Message::HEADERFLAG_QR); + message.clear(Message::RENDER); + message.setQid(qid); + message.setOpcode(opcode); + message.setHeaderFlag(Message::HEADERFLAG_QR); if (rd) { - message->setHeaderFlag(Message::HEADERFLAG_RD); + message.setHeaderFlag(Message::HEADERFLAG_RD); } if (cd) { - message->setHeaderFlag(Message::HEADERFLAG_CD); + message.setHeaderFlag(Message::HEADERFLAG_CD); } for_each(questions.begin(), questions.end(), QuestionInserter(message)); - message->setRcode(rcode); + message.setRcode(rcode); - MessageRenderer renderer(buffer); + MessageRenderer renderer; - renderer.setBuffer(buffer.get()); ++ renderer.setBuffer(&buffer); if (tsig_context.get() != NULL) { - message->toWire(renderer, *tsig_context); + message.toWire(renderer, *tsig_context); } else { - message->toWire(renderer); + message.toWire(renderer); } + renderer.setBuffer(NULL); LOG_DEBUG(auth_logger, DBG_AUTH_MESSAGES, AUTH_SEND_ERROR_RESPONSE) - .arg(renderer.getLength()).arg(*message); + .arg(renderer.getLength()).arg(message); } } @@@ -481,43 -478,35 +480,43 @@@ AuthSrv::processMessage(const IOMessage return; } - // update per opcode statistics counter. This can only be reliable after - // TSIG check succeeds. - impl_->counters_.inc(message.getOpcode()); - bool send_answer = true; - if (message.getOpcode() == Opcode::NOTIFY()) { - send_answer = impl_->processNotify(io_message, message, buffer, - tsig_context); - } else if (message.getOpcode() != Opcode::QUERY()) { - LOG_DEBUG(auth_logger, DBG_AUTH_DETAIL, AUTH_UNSUPPORTED_OPCODE) - .arg(message.getOpcode().toText()); - makeErrorMessage(message, buffer, Rcode::NOTIMP(), tsig_context); - } else if (message.getRRCount(Message::SECTION_QUESTION) != 1) { - makeErrorMessage(message, buffer, Rcode::FORMERR(), tsig_context); - } else { - ConstQuestionPtr question = *message.beginQuestion(); - const RRType &qtype = question->getType(); - if (qtype == RRType::AXFR()) { - send_answer = impl_->processXfrQuery(io_message, message, buffer, - tsig_context); - } else if (qtype == RRType::IXFR()) { - send_answer = impl_->processXfrQuery(io_message, message, buffer, - tsig_context); + try { + // update per opcode statistics counter. This can only be reliable + // after TSIG check succeeds. - impl_->counters_.inc(message->getOpcode()); ++ impl_->counters_.inc(message.getOpcode()); + - if (message->getOpcode() == Opcode::NOTIFY()) { ++ if (message.getOpcode() == Opcode::NOTIFY()) { + send_answer = impl_->processNotify(io_message, message, buffer, + tsig_context); - } else if (message->getOpcode() != Opcode::QUERY()) { ++ } else if (message.getOpcode() != Opcode::QUERY()) { + LOG_DEBUG(auth_logger, DBG_AUTH_DETAIL, AUTH_UNSUPPORTED_OPCODE) - .arg(message->getOpcode().toText()); ++ .arg(message.getOpcode().toText()); + makeErrorMessage(message, buffer, Rcode::NOTIMP(), tsig_context); - } else if (message->getRRCount(Message::SECTION_QUESTION) != 1) { ++ } else if (message.getRRCount(Message::SECTION_QUESTION) != 1) { + makeErrorMessage(message, buffer, Rcode::FORMERR(), tsig_context); } else { - ConstQuestionPtr question = *message->beginQuestion(); - send_answer = impl_->processNormalQuery(io_message, message, - buffer, tsig_context); ++ ConstQuestionPtr question = *message.beginQuestion(); + const RRType &qtype = question->getType(); + if (qtype == RRType::AXFR()) { + send_answer = impl_->processXfrQuery(io_message, message, + buffer, tsig_context); + } else if (qtype == RRType::IXFR()) { + send_answer = impl_->processXfrQuery(io_message, message, + buffer, tsig_context); + } else { + send_answer = impl_->processNormalQuery(io_message, message, + buffer, tsig_context); + } } + } catch (const std::exception& ex) { + LOG_DEBUG(auth_logger, DBG_AUTH_DETAIL, AUTH_RESPONSE_FAILURE) + .arg(ex.what()); + makeErrorMessage(message, buffer, Rcode::SERVFAIL()); + } catch (...) { + LOG_DEBUG(auth_logger, DBG_AUTH_DETAIL, AUTH_RESPONSE_FAILURE_UNKNOWN); + makeErrorMessage(message, buffer, Rcode::SERVFAIL()); } - impl_->resumeServer(server, message, send_answer); } @@@ -564,19 -553,17 +563,19 @@@ AuthSrvImpl::processNormalQuery(const I return (true); } - MessageRenderer renderer(buffer); + MessageRenderer renderer; - renderer.setBuffer(buffer.get()); ++ renderer.setBuffer(&buffer); const bool udp_buffer = (io_message.getSocket().getProtocol() == IPPROTO_UDP); renderer.setLengthLimit(udp_buffer ? remote_bufsize : 65535); if (tsig_context.get() != NULL) { - message->toWire(renderer, *tsig_context); + message.toWire(renderer, *tsig_context); } else { - message->toWire(renderer); + message.toWire(renderer); } + renderer.setBuffer(NULL); LOG_DEBUG(auth_logger, DBG_AUTH_MESSAGES, AUTH_SEND_NORMAL_RESPONSE) - .arg(renderer.getLength()).arg(message->toText()); + .arg(renderer.getLength()).arg(message); return (true); } @@@ -691,18 -678,16 +690,18 @@@ AuthSrvImpl::processNotify(const IOMess return (false); } - message->makeResponse(); - message->setHeaderFlag(Message::HEADERFLAG_AA); - message->setRcode(Rcode::NOERROR()); + message.makeResponse(); + message.setHeaderFlag(Message::HEADERFLAG_AA); + message.setRcode(Rcode::NOERROR()); - MessageRenderer renderer(buffer); + MessageRenderer renderer; - renderer.setBuffer(buffer.get()); ++ renderer.setBuffer(&buffer); if (tsig_context.get() != NULL) { - message->toWire(renderer, *tsig_context); + message.toWire(renderer, *tsig_context); } else { - message->toWire(renderer); + message.toWire(renderer); } + renderer.setBuffer(NULL); return (true); } diff --cc src/bin/auth/tests/auth_srv_unittest.cc index c742c54727,862fef5874..688ce62a03 --- a/src/bin/auth/tests/auth_srv_unittest.cc +++ b/src/bin/auth/tests/auth_srv_unittest.cc @@@ -87,12 -87,8 +87,12 @@@ protected server.setXfrinSession(¬ify_session); server.setStatisticsSession(&statistics_session); } + virtual void processMessage() { + // If processMessage has been called before, parse_message needs + // to be reset. If it hasn't, there's no harm in doing so + parse_message->clear(Message::PARSE); - server.processMessage(*io_message, parse_message, response_obuffer, + server.processMessage(*io_message, *parse_message, *response_obuffer, &dnsserv); } @@@ -502,9 -492,9 +506,9 @@@ TEST_F(AuthSrvTest, AXFRDisconnectFail Name("example.com"), RRClass::IN(), RRType::AXFR()); createRequestPacket(request_message, IPPROTO_TCP); - EXPECT_NO_THROW(server.processMessage(*io_message, parse_message, - response_obuffer, &dnsserv)); - EXPECT_THROW(server.processMessage(*io_message, *parse_message, - *response_obuffer, &dnsserv), - XfroutError); ++ EXPECT_NO_THROW(server.processMessage(*io_message, *parse_message, ++ *response_obuffer, &dnsserv)); + // Since the disconnect failed, we should still be 'connected' EXPECT_TRUE(xfrout.isConnected()); // XXX: we need to re-enable disconnect. otherwise an exception would be // thrown via the destructor of the server. @@@ -560,8 -553,9 +567,8 @@@ TEST_F(AuthSrvTest, IXFRDisconnectFail Name("example.com"), RRClass::IN(), RRType::IXFR()); createRequestPacket(request_message, IPPROTO_TCP); - EXPECT_NO_THROW(server.processMessage(*io_message, parse_message, - response_obuffer, &dnsserv)); - EXPECT_THROW(server.processMessage(*io_message, *parse_message, - *response_obuffer, &dnsserv), - XfroutError); ++ EXPECT_NO_THROW(server.processMessage(*io_message, *parse_message, ++ *response_obuffer, &dnsserv)); EXPECT_TRUE(xfrout.isConnected()); // XXX: we need to re-enable disconnect. otherwise an exception would be // thrown via the destructor of the server. @@@ -1050,231 -1064,4 +1075,231 @@@ TEST_F(AuthSrvTest, listenAddresses) "Released tokens"); } +// +// Tests for catching exceptions in various stages of the query processing +// +// These tests work by defining two proxy classes, that act as an in-memory +// client by default, but can throw exceptions at various points. +// +namespace { + +/// A the possible methods to throw in, either in FakeInMemoryClient or +/// FakeZoneFinder +enum ThrowWhen { + THROW_NEVER, + THROW_AT_FIND_ZONE, + THROW_AT_GET_ORIGIN, + THROW_AT_GET_CLASS, + THROW_AT_FIND, + THROW_AT_FIND_ALL, + THROW_AT_FIND_NSEC3 +}; + +/// convenience function to check whether and what to throw +void +checkThrow(ThrowWhen method, ThrowWhen throw_at, bool isc_exception) { + if (method == throw_at) { + if (isc_exception) { + isc_throw(isc::Exception, "foo"); + } else { + throw std::exception(); + } + } +} + +/// \brief proxy class for the ZoneFinder returned by the InMemoryClient +/// proxied by FakeInMemoryClient +/// +/// See the documentation for FakeInMemoryClient for more information, +/// all methods simply check whether they should throw, and if not, call +/// their proxied equivalent. +class FakeZoneFinder : public isc::datasrc::ZoneFinder { +public: + FakeZoneFinder(isc::datasrc::ZoneFinderPtr zone_finder, + ThrowWhen throw_when, + bool isc_exception) : + real_zone_finder_(zone_finder), + throw_when_(throw_when), + isc_exception_(isc_exception) + {} + + virtual isc::dns::Name + getOrigin() const { + checkThrow(THROW_AT_GET_ORIGIN, throw_when_, isc_exception_); + return (real_zone_finder_->getOrigin()); + } + + virtual isc::dns::RRClass + getClass() const { + checkThrow(THROW_AT_GET_CLASS, throw_when_, isc_exception_); + return (real_zone_finder_->getClass()); + } + + virtual isc::datasrc::ZoneFinder::FindResult + find(const isc::dns::Name& name, + const isc::dns::RRType& type, + isc::datasrc::ZoneFinder::FindOptions options) + { + checkThrow(THROW_AT_FIND, throw_when_, isc_exception_); + return (real_zone_finder_->find(name, type, options)); + } + + virtual FindResult + findAll(const isc::dns::Name& name, + std::vector &target, + const FindOptions options = FIND_DEFAULT) + { + checkThrow(THROW_AT_FIND_ALL, throw_when_, isc_exception_); + return (real_zone_finder_->findAll(name, target, options)); + }; + + virtual FindNSEC3Result + findNSEC3(const isc::dns::Name& name, bool recursive) { + checkThrow(THROW_AT_FIND_NSEC3, throw_when_, isc_exception_); + return (real_zone_finder_->findNSEC3(name, recursive)); + }; + + virtual isc::dns::Name + findPreviousName(const isc::dns::Name& query) const { + return (real_zone_finder_->findPreviousName(query)); + } + +private: + isc::datasrc::ZoneFinderPtr real_zone_finder_; + ThrowWhen throw_when_; + bool isc_exception_; +}; + +/// \brief Proxy InMemoryClient that can throw exceptions at specified times +/// +/// It is based on the memory client since that one is easy to override +/// (with setInMemoryClient) with the current design of AuthSrv. +class FakeInMemoryClient : public isc::datasrc::InMemoryClient { +public: + /// \brief Create a proxy memory client + /// + /// \param real_client The real in-memory client to proxy + /// \param throw_when if set to any value other than never, that is + /// the method that will throw an exception (either in this + /// class or the related FakeZoneFinder) + /// \param isc_exception if true, throw isc::Exception, otherwise, + /// throw std::exception + FakeInMemoryClient(AuthSrv::InMemoryClientPtr real_client, + ThrowWhen throw_when, + bool isc_exception) : + real_client_(real_client), + throw_when_(throw_when), + isc_exception_(isc_exception) + {} + + /// \brief proxy call for findZone + /// + /// if this instance was constructed with throw_when set to find_zone, + /// this method will throw. Otherwise, it will return a FakeZoneFinder + /// instance which will throw at the method specified at the + /// construction of this instance. + virtual FindResult + findZone(const isc::dns::Name& name) const { + checkThrow(THROW_AT_FIND_ZONE, throw_when_, isc_exception_); + const FindResult result = real_client_->findZone(name); + return (FindResult(result.code, isc::datasrc::ZoneFinderPtr( + new FakeZoneFinder(result.zone_finder, + throw_when_, + isc_exception_)))); + } + +private: + AuthSrv::InMemoryClientPtr real_client_; + ThrowWhen throw_when_; + bool isc_exception_; +}; + +} // end anonymous namespace for throwing proxy classes + +// Test for the tests +// +// Set the proxies to never throw, this should have the same result as +// queryWithInMemoryClientNoDNSSEC, and serves to test the two proxy classes +TEST_F(AuthSrvTest, queryWithInMemoryClientProxy) { + // Set real inmem client to proxy + updateConfig(&server, CONFIG_INMEMORY_EXAMPLE, true); + + AuthSrv::InMemoryClientPtr fake_client( + new FakeInMemoryClient(server.getInMemoryClient(rrclass), + THROW_NEVER, + false)); + + ASSERT_NE(AuthSrv::InMemoryClientPtr(), server.getInMemoryClient(rrclass)); + server.setInMemoryClient(rrclass, fake_client); + + createDataFromFile("nsec3query_nodnssec_fromWire.wire"); - server.processMessage(*io_message, parse_message, response_obuffer, ++ server.processMessage(*io_message, *parse_message, *response_obuffer, + &dnsserv); + + EXPECT_TRUE(dnsserv.hasAnswer()); + headerCheck(*parse_message, default_qid, Rcode::NOERROR(), + opcode.getCode(), QR_FLAG | AA_FLAG, 1, 1, 2, 1); +} + +// Convenience function for the rest of the tests, set up a proxy +// to throw in the given method +// If isc_exception is true, it will throw isc::Exception, otherwise +// it will throw std::exception +void +setupThrow(AuthSrv* server, const char *config, ThrowWhen throw_when, + bool isc_exception) +{ + // Set real inmem client to proxy + updateConfig(server, config, true); + + // Set it to throw on findZone(), this should result in + // SERVFAIL on any exception + AuthSrv::InMemoryClientPtr fake_client( + new FakeInMemoryClient( + server->getInMemoryClient(isc::dns::RRClass::IN()), + throw_when, + isc_exception)); + + ASSERT_NE(AuthSrv::InMemoryClientPtr(), + server->getInMemoryClient(isc::dns::RRClass::IN())); + server->setInMemoryClient(isc::dns::RRClass::IN(), fake_client); +} + +TEST_F(AuthSrvTest, queryWithThrowingProxyServfails) { + // Test the common cases, all of which should simply return SERVFAIL + // Use THROW_NEVER as end marker + ThrowWhen throws[] = { THROW_AT_FIND_ZONE, + THROW_AT_GET_ORIGIN, + THROW_AT_FIND, + THROW_AT_FIND_NSEC3, + THROW_NEVER }; + UnitTestUtil::createDNSSECRequestMessage(request_message, opcode, + default_qid, Name("foo.example."), + RRClass::IN(), RRType::TXT()); + for (ThrowWhen* when(throws); *when != THROW_NEVER; ++when) { + createRequestPacket(request_message, IPPROTO_UDP); + setupThrow(&server, CONFIG_INMEMORY_EXAMPLE, *when, true); + processAndCheckSERVFAIL(); + // To be sure, check same for non-isc-exceptions + createRequestPacket(request_message, IPPROTO_UDP); + setupThrow(&server, CONFIG_INMEMORY_EXAMPLE, *when, false); + processAndCheckSERVFAIL(); + } +} + +// Throw isc::Exception in getClass(). (Currently?) getClass is not called +// in the processMessage path, so this should result in a normal answer +TEST_F(AuthSrvTest, queryWithInMemoryClientProxyGetClass) { + createDataFromFile("nsec3query_nodnssec_fromWire.wire"); + setupThrow(&server, CONFIG_INMEMORY_EXAMPLE, THROW_AT_GET_CLASS, true); + + // getClass is not called so it should just answer - server.processMessage(*io_message, parse_message, response_obuffer, ++ server.processMessage(*io_message, *parse_message, *response_obuffer, + &dnsserv); + + EXPECT_TRUE(dnsserv.hasAnswer()); + headerCheck(*parse_message, default_qid, Rcode::NOERROR(), + opcode.getCode(), QR_FLAG | AA_FLAG, 1, 1, 2, 1); +} + }